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 ffmpeg
import pathlib
import queue
import shutil
import threading
import numpy as np
from queue import Queue
from shutil import rmtree
from PIL import Image
from rembg.bg import remove
from rembg import new_session
from rembg import new_session, remove
os.environ.setdefault(
"PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True,managed_memory:True"
@@ -25,14 +24,14 @@ parser.add_argument(
parser.add_argument(
"--model",
type=str,
default="birefnet-general-lite",
default="u2net-human-seg",
help="rembg model to use (default: birefnet-general-lite)",
)
parser.add_argument(
"--workers",
type=int,
default=4,
help="Number of concurrent processing workers (default: 4)",
default=1,
help="Number of concurrent processing workers (default: 1)",
)
parser.add_argument(
"--smooth",
@@ -52,12 +51,6 @@ parser.add_argument(
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(
"--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")
# Extract input video frames
if not args.skip_extract:
if not os.path.isdir(frames_dir):
os.mkdir(frames_dir)
stream = ffmpeg.input(args.input)
@@ -88,7 +80,6 @@ _SENTINEL = object()
try:
# Process frames with pipelined reader -> processors -> writer
if not args.skip_process:
if not os.path.isdir(processed_dir):
os.mkdir(processed_dir)
@@ -96,10 +87,10 @@ try:
total_files = len(files)
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)
write_queue = queue.Queue(maxsize=args.write_buffer)
read_queue = Queue(maxsize=args.read_ahead)
write_queue = Queue(maxsize=args.write_buffer)
errors = []
active_processors = [
args.workers
@@ -170,7 +161,7 @@ try:
if errors:
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:
files = sorted(os.listdir(processed_dir))
total = len(files)
@@ -181,13 +172,11 @@ try:
buf = {} # read_idx -> (filename, np.ndarray RGBA)
for read_idx in range(total + half):
# Load next frame into buffer
if read_idx < total:
file = files[read_idx]
img = Image.open(os.path.join(processed_dir, file)).convert("RGBA")
buf[read_idx] = (file, np.array(img))
# The frame we can now finalize (has full right-side context)
write_idx = read_idx - half
if 0 <= write_idx < total:
start = max(0, write_idx - half)
@@ -205,7 +194,6 @@ try:
Image.fromarray(out_arr).save(os.path.join(processed_dir, filename))
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
if drop_idx in buf:
del buf[drop_idx]
@@ -231,5 +219,5 @@ except KeyboardInterrupt:
finally:
print("Removing temporary files...")
shutil.rmtree(processed_dir, ignore_errors=True)
shutil.rmtree(frames_dir, ignore_errors=True)
rmtree(processed_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