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
+84 -96
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,100 +70,98 @@ 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) stream = ffmpeg.output(stream, os.path.join(frames_dir, "%04d.png"))
stream = ffmpeg.output(stream, os.path.join(frames_dir, "%04d.png")) ffmpeg.run(stream)
ffmpeg.run(stream)
_SENTINEL = object() _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)
files = sorted(os.listdir(frames_dir)) files = sorted(os.listdir(frames_dir))
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
] # list so processor() can mutate without nonlocal ] # list so processor() can mutate without nonlocal
active_processors_lock = threading.Lock() active_processors_lock = threading.Lock()
def reader(): def reader():
try: try:
for idx, file in enumerate(files, 1): for idx, file in enumerate(files, 1):
with open(os.path.join(frames_dir, file), "rb") as f: with open(os.path.join(frames_dir, file), "rb") as f:
data = f.read() data = f.read()
read_queue.put((idx, file, data)) read_queue.put((idx, file, data))
except Exception as e: except Exception as e:
errors.append(e) errors.append(e)
finally: finally:
# One sentinel per worker so each one knows when to stop # One sentinel per worker so each one knows when to stop
for _ in range(args.workers): for _ in range(args.workers):
read_queue.put(_SENTINEL) read_queue.put(_SENTINEL)
def processor(): def processor():
try: try:
while True: while True:
item = read_queue.get() item = read_queue.get()
if item is _SENTINEL: if item is _SENTINEL:
break break
idx, file, input_data = item idx, file, input_data = item
print(f"Processing frame {idx}/{total_files}: {file}", flush=True) print(f"Processing frame {idx}/{total_files}: {file}", flush=True)
output_data = remove(input_data, session=session) output_data = remove(input_data, session=session)
write_queue.put((idx, file, output_data)) write_queue.put((idx, file, output_data))
except Exception as e: except Exception as e:
errors.append(e) errors.append(e)
finally: finally:
# Signal writer only when the last processor finishes # Signal writer only when the last processor finishes
with active_processors_lock: with active_processors_lock:
active_processors[0] -= 1 active_processors[0] -= 1
if active_processors[0] == 0: if active_processors[0] == 0:
write_queue.put(_SENTINEL) write_queue.put(_SENTINEL)
def writer(): def writer():
try: try:
while True: while True:
item = write_queue.get() item = write_queue.get()
if item is _SENTINEL: if item is _SENTINEL:
break break
idx, file, output_data = item idx, file, output_data = item
with open(os.path.join(processed_dir, file), "wb") as f: with open(os.path.join(processed_dir, file), "wb") as f:
f.write(output_data) f.write(output_data)
print(f"Written frame {idx}/{total_files}: {file}", flush=True) print(f"Written frame {idx}/{total_files}: {file}", flush=True)
except Exception as e: except Exception as e:
errors.append(e) errors.append(e)
reader_thread = threading.Thread(target=reader, daemon=True) reader_thread = threading.Thread(target=reader, daemon=True)
processor_threads = [ processor_threads = [
threading.Thread(target=processor, daemon=True) for _ in range(args.workers) threading.Thread(target=processor, daemon=True) for _ in range(args.workers)
] ]
writer_thread = threading.Thread(target=writer, daemon=True) writer_thread = threading.Thread(target=writer, daemon=True)
reader_thread.start() reader_thread.start()
for t in processor_threads: for t in processor_threads:
t.start() t.start()
writer_thread.start() writer_thread.start()
reader_thread.join() reader_thread.join()
for t in processor_threads: for t in processor_threads:
t.join() t.join()
writer_thread.join() writer_thread.join()
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