Fixed CUDA version mismatch

This commit is contained in:
2026-06-23 14:47:34 +03:00
parent 1b5f92ccb1
commit a8f92251a3
2 changed files with 126 additions and 96 deletions
+12 -24
View File
@@ -2,13 +2,12 @@ import argparse
import os import os
import ffmpeg import ffmpeg
import pathlib import pathlib
import queue
import shutil
import threading import threading
import numpy as np import numpy as np
from queue import Queue
from shutil import rmtree
from PIL import Image from PIL import Image
from rembg.bg import remove from rembg import new_session, remove
from rembg import new_session
os.environ.setdefault( os.environ.setdefault(
"PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True,managed_memory:True" "PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True,managed_memory:True"
@@ -25,14 +24,14 @@ parser.add_argument(
parser.add_argument( parser.add_argument(
"--model", "--model",
type=str, type=str,
default="birefnet-general-lite", default="u2net-human-seg",
help="rembg model to use (default: birefnet-general-lite)", help="rembg model to use (default: birefnet-general-lite)",
) )
parser.add_argument( parser.add_argument(
"--workers", "--workers",
type=int, type=int,
default=4, default=1,
help="Number of concurrent processing workers (default: 4)", help="Number of concurrent processing workers (default: 1)",
) )
parser.add_argument( parser.add_argument(
"--smooth", "--smooth",
@@ -52,12 +51,6 @@ parser.add_argument(
default=8, default=8,
help="Number of processed frames to buffer before writing (default: 8)", help="Number of processed frames to buffer before writing (default: 8)",
) )
parser.add_argument(
"--skip-extract", action="store_true", help="Skips ffmpeg frame extraction"
)
parser.add_argument(
"--skip-process", action="store_true", help="Skips rembg frame processing"
)
parser.add_argument( parser.add_argument(
"--skip-smooth", action="store_true", help="Skips temporal mask smoothing" "--skip-smooth", action="store_true", help="Skips temporal mask smoothing"
) )
@@ -77,7 +70,6 @@ frames_dir = os.path.join(str(pathlib.Path(__file__).parent.absolute()), "frames
processed_dir = os.path.join(str(pathlib.Path(__file__).parent.absolute()), "processed") processed_dir = os.path.join(str(pathlib.Path(__file__).parent.absolute()), "processed")
# Extract input video frames # Extract input video frames
if not args.skip_extract:
if not os.path.isdir(frames_dir): if not os.path.isdir(frames_dir):
os.mkdir(frames_dir) os.mkdir(frames_dir)
stream = ffmpeg.input(args.input) stream = ffmpeg.input(args.input)
@@ -88,7 +80,6 @@ _SENTINEL = object()
try: try:
# Process frames with pipelined reader -> processors -> writer # Process frames with pipelined reader -> processors -> writer
if not args.skip_process:
if not os.path.isdir(processed_dir): if not os.path.isdir(processed_dir):
os.mkdir(processed_dir) os.mkdir(processed_dir)
@@ -96,10 +87,10 @@ try:
total_files = len(files) total_files = len(files)
print(f"Loading rembg session (model={args.model})...", flush=True) print(f"Loading rembg session (model={args.model})...", flush=True)
session = new_session(args.model) session = new_session(args.model, providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
read_queue = queue.Queue(maxsize=args.read_ahead) read_queue = Queue(maxsize=args.read_ahead)
write_queue = queue.Queue(maxsize=args.write_buffer) write_queue = Queue(maxsize=args.write_buffer)
errors = [] errors = []
active_processors = [ active_processors = [
args.workers args.workers
@@ -170,7 +161,7 @@ try:
if errors: if errors:
raise errors[0] raise errors[0]
# Temporal mask smoothing — streaming sliding window, only `window` frames in RAM at once # Temporal mask smoothing
if not args.skip_smooth and args.smooth > 0: if not args.skip_smooth and args.smooth > 0:
files = sorted(os.listdir(processed_dir)) files = sorted(os.listdir(processed_dir))
total = len(files) total = len(files)
@@ -181,13 +172,11 @@ try:
buf = {} # read_idx -> (filename, np.ndarray RGBA) buf = {} # read_idx -> (filename, np.ndarray RGBA)
for read_idx in range(total + half): for read_idx in range(total + half):
# Load next frame into buffer
if read_idx < total: if read_idx < total:
file = files[read_idx] file = files[read_idx]
img = Image.open(os.path.join(processed_dir, file)).convert("RGBA") img = Image.open(os.path.join(processed_dir, file)).convert("RGBA")
buf[read_idx] = (file, np.array(img)) buf[read_idx] = (file, np.array(img))
# The frame we can now finalize (has full right-side context)
write_idx = read_idx - half write_idx = read_idx - half
if 0 <= write_idx < total: if 0 <= write_idx < total:
start = max(0, write_idx - half) start = max(0, write_idx - half)
@@ -205,7 +194,6 @@ try:
Image.fromarray(out_arr).save(os.path.join(processed_dir, filename)) Image.fromarray(out_arr).save(os.path.join(processed_dir, filename))
print(f"Smoothed frame {write_idx + 1}/{total}: {filename}", flush=True) print(f"Smoothed frame {write_idx + 1}/{total}: {filename}", flush=True)
# Drop the frame that's no longer needed by any future window
drop_idx = write_idx - half drop_idx = write_idx - half
if drop_idx in buf: if drop_idx in buf:
del buf[drop_idx] del buf[drop_idx]
@@ -231,5 +219,5 @@ except KeyboardInterrupt:
finally: finally:
print("Removing temporary files...") print("Removing temporary files...")
shutil.rmtree(processed_dir, ignore_errors=True) rmtree(processed_dir, ignore_errors=True)
shutil.rmtree(frames_dir, ignore_errors=True) rmtree(frames_dir, ignore_errors=True)
+42
View File
@@ -0,0 +1,42 @@
attrs==26.1.0
certifi==2026.6.17
charset-normalizer==3.4.7
coloredlogs==15.0.1
ffmpeg-python==0.2.0
flatbuffers==25.12.19
future==1.0.0
humanfriendly==10.0
idna==3.18
ImageIO==2.37.3
jsonschema==4.26.0
jsonschema-specifications==2025.9.1
lazy-loader==0.5
llvmlite==0.47.0
mpmath==1.3.0
networkx==3.6.1
numba==0.65.1
numpy==2.4.6
nvidia-cublas-cu12==12.9.2.10
nvidia-cuda-nvrtc-cu12==12.9.86
nvidia-cuda-runtime-cu12==12.9.79
nvidia-cudnn-cu12==9.10.2.21
nvidia-cufft-cu12==11.4.1.4
nvidia-nvjitlink-cu12==12.9.86
onnxruntime-gpu==1.19.2
packaging==26.2
pillow==12.2.0
platformdirs==4.10.0
pooch==1.9.0
protobuf==7.35.1
PyMatting==1.1.15
referencing==0.37.0
rembg==2.0.76
requests==2.34.2
rpds-py==2026.5.1
scikit-image==0.26.0
scipy==1.18.0
sympy==1.14.0
tifffile==2026.6.1
tqdm==4.68.3
typing_extensions==4.15.0
urllib3==2.7.0