aboutsummaryrefslogtreecommitdiff
path: root/lib/project/src/voiceforge/ai_worker.py
diff options
context:
space:
mode:
Diffstat (limited to 'lib/project/src/voiceforge/ai_worker.py')
-rw-r--r--lib/project/src/voiceforge/ai_worker.py161
1 files changed, 161 insertions, 0 deletions
diff --git a/lib/project/src/voiceforge/ai_worker.py b/lib/project/src/voiceforge/ai_worker.py
new file mode 100644
index 0000000..99ec9f5
--- /dev/null
+++ b/lib/project/src/voiceforge/ai_worker.py
@@ -0,0 +1,161 @@
+"""Private script executed only by the isolated backend Python."""
+from __future__ import annotations
+
+import argparse
+import hashlib
+import json
+import os
+from pathlib import Path
+import signal
+import tempfile
+import urllib.request
+import zipfile
+
+
+def report(stage: str, completed: float | None = None, total: float | None = None) -> None:
+ print("VOICEFORGE_PROGRESS " + json.dumps([stage, completed, total]), flush=True)
+
+
+def ensure_model(directory: Path) -> Path:
+ """Install a release-pinned, hash-checked archive atomically under setup lock."""
+ model = directory / "DeepFilterNet3"
+ if (model / "config.ini").is_file() and (model / "checkpoints/model_120.ckpt.best").is_file():
+ return model
+ if directory.exists():
+ raise RuntimeError(f"Incomplete model directory: {directory}. Remove it and retry AI setup.")
+ directory.parent.mkdir(parents=True, exist_ok=True)
+ url = "https://raw.githubusercontent.com/Rikorose/DeepFilterNet/v0.5.6/models/DeepFilterNet3.zip"
+ expected = "49c52edc8947ae1f9bf50d81530beaf3a2c3245aeaf34b6f31ff535cd22284d2"
+ # Release archive size is fixed alongside its hash, even without Content-Length.
+ expected_bytes = 7986207
+ downloaded = 0
+ report("Downloading DeepFilterNet3 model (bytes)", 0, expected_bytes)
+ with tempfile.TemporaryDirectory(prefix=".df3-", dir=directory.parent) as staging:
+ archive = Path(staging) / "model.zip"
+ digest = hashlib.sha256()
+ with urllib.request.urlopen(url, timeout=60) as response, archive.open("wb") as output:
+ while chunk := response.read(256 * 1024):
+ downloaded += len(chunk)
+ if downloaded > expected_bytes:
+ raise RuntimeError("DeepFilterNet3 model download exceeds the expected size")
+ digest.update(chunk)
+ output.write(chunk)
+ report("Downloading DeepFilterNet3 model (bytes)", downloaded, expected_bytes)
+ if downloaded != expected_bytes or digest.hexdigest() != expected:
+ raise RuntimeError("DeepFilterNet3 model checksum mismatch; refusing to load weights")
+ report("Verifying and extracting DeepFilterNet3 model")
+ unpacked = Path(staging) / "unpacked"
+ with zipfile.ZipFile(archive) as bundle:
+ for member in bundle.infolist():
+ destination = (unpacked / member.filename).resolve()
+ if not destination.is_relative_to(unpacked.resolve()):
+ raise RuntimeError("Unsafe path in model archive")
+ bundle.extractall(unpacked)
+ os.replace(unpacked, directory)
+ return model
+
+
+def main() -> None:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--device", choices=("auto", "cpu", "cuda"), required=True)
+ parser.add_argument("--check", action="store_true")
+ parser.add_argument("--model-dir", type=Path, required=True)
+ parser.add_argument("--source", type=Path)
+ parser.add_argument("--target", type=Path)
+ parser.add_argument("--strength", type=float, default=1.0)
+ args = parser.parse_args()
+
+ # SIGTERM from the parent takes the same cleanup path as Ctrl-C.
+ def terminate(signum, frame):
+ raise KeyboardInterrupt
+
+ signal.signal(signal.SIGTERM, terminate)
+
+ import numpy as np
+ import soundfile as sf
+ import torch
+
+ torch.set_num_threads(min(4, os.cpu_count() or 1))
+ available = torch.cuda.is_available()
+ if args.device == "cuda" and not available:
+ raise RuntimeError("CUDA was explicitly requested, but PyTorch cannot use it. "
+ "Check the NVIDIA driver and CUDA_VISIBLE_DEVICES, or select cpu.")
+ device = "cuda" if args.device != "cpu" and available else "cpu"
+ # df.utils.get_device() consults DEVICE on every call, including model creation.
+ os.environ["DEVICE"] = device
+ from df.enhance import enhance, init_df
+ from df.model import ModelParams
+ from libdf import DF
+
+ model_path = ensure_model(args.model_dir)
+ report(f"Loading DeepFilterNet3 ({device})")
+ model, state, _ = init_df(str(model_path), log_level="ERROR", log_file=None)
+ if state.sr() != 48000:
+ raise RuntimeError("DeepFilterNet3 model must use 48 kHz")
+ if args.check:
+ # Exercise both model execution and the compiled libdf/NumPy ABI.
+ sample = enhance(model, state, torch.zeros(1, 4800), pad=True)
+ if sample.shape != (1, 4800) or not torch.isfinite(sample).all():
+ raise RuntimeError("DeepFilterNet3 self-test failed")
+ report(f"DeepFilterNet3 ready ({device})", 1, 1)
+ return
+ if args.source is None or args.target is None:
+ parser.error("--source and --target are required unless --check is used")
+ if not np.isfinite(args.strength) or not 0 <= args.strength <= 1:
+ parser.error("--strength must be in [0, 1]")
+
+ temporary = None
+ try:
+ with sf.SoundFile(args.source) as source:
+ if source.samplerate != 48000 or source.channels != 1 or source.format not in {"WAV", "WAVEX", "RF64"}:
+ raise ValueError("AI input must be a mono 48 kHz WAV; convert it before denoising")
+ frames = len(source)
+ fd, name = tempfile.mkstemp(prefix=".voiceforge-ai-", suffix=".wav", dir=args.target.parent)
+ os.close(fd)
+ temporary = Path(name)
+ # RF64 avoids RIFF's 4 GiB limit on very long recordings.
+ output_format = "RF64" if frames * 4 > 0xFFFFFFFF - 4096 else "WAV"
+ with sf.SoundFile(temporary, "w", samplerate=48000, channels=1,
+ format=output_format, subtype="FLOAT") as target:
+ block, context, overlap = 480000, 96000, 2400
+ previous = None
+ params = ModelParams()
+ report(f"Denoising ({device})", 0, frames)
+ for start in range(0, frames, block):
+ end = min(start + block, frames)
+ left, right = max(0, start - context), min(frames, end + overlap + context)
+ source.seek(left)
+ audio = source.read(right - left, dtype="float32")
+ if not np.isfinite(audio).all():
+ raise ValueError("Input contains non-finite audio samples")
+ if args.strength:
+ # Fresh STFT state per window; context warms normalization and
+ # recurrent layers. Pad to a whole hop before delay compensation.
+ state = DF(sr=params.sr, fft_size=params.fft_size,
+ hop_size=params.hop_size, nb_bands=params.nb_erb,
+ min_nb_erb_freqs=params.min_nb_freqs)
+ padded = np.pad(audio, (0, (-len(audio)) % params.hop_size))
+ wet = enhance(model, state, torch.from_numpy(padded).unsqueeze(0),
+ pad=True).squeeze(0).numpy()[:len(audio)]
+ if len(wet) != len(audio) or not np.isfinite(wet).all():
+ raise RuntimeError("DeepFilterNet3 returned invalid audio")
+ audio = audio * (1 - args.strength) + wet * args.strength
+ kept = audio[start - left:min(frames, end + overlap) - left].copy()
+ if previous is not None and args.strength:
+ n = len(previous)
+ fade = np.linspace(0, 1, n, dtype=np.float32)
+ kept[:n] = previous * (1 - fade) + kept[:n] * fade
+ target.write(kept[:end - start])
+ previous = kept[end - start:].copy()
+ report(f"Denoising ({device})", end, frames)
+ if target.tell() != frames:
+ raise RuntimeError("Output frame count does not match input")
+ os.replace(temporary, args.target)
+ temporary = None
+ finally:
+ if temporary is not None:
+ temporary.unlink(missing_ok=True)
+
+
+if __name__ == "__main__":
+ main()