41 lines
1.3 KiB
Python
41 lines
1.3 KiB
Python
"""Remove training-only imports from Resemble Enhance's inference path."""
|
|
|
|
from pathlib import Path
|
|
|
|
import resemble_enhance
|
|
|
|
|
|
root = Path(resemble_enhance.__file__).parent
|
|
|
|
replacements = {
|
|
root / "enhancer" / "inference.py": {
|
|
"from .train import Enhancer, HParams": (
|
|
"from .enhancer import Enhancer\nfrom .hparams import HParams"
|
|
),
|
|
},
|
|
root / "denoiser" / "inference.py": {
|
|
"from .train import Denoiser, HParams": (
|
|
"from .denoiser import Denoiser\nfrom .hparams import HParams"
|
|
),
|
|
},
|
|
root / "enhancer" / "enhancer.py": {
|
|
"from ..utils.distributed import global_leader_only\n"
|
|
"from ..utils.train_loop import TrainLoop": (
|
|
"def global_leader_only(fn):\n"
|
|
" return fn\n\n"
|
|
"class TrainLoop:\n"
|
|
" @classmethod\n"
|
|
" def get_running_loop(cls):\n"
|
|
" return None"
|
|
),
|
|
},
|
|
}
|
|
|
|
for path, edits in replacements.items():
|
|
text = path.read_text(encoding="utf-8")
|
|
for old, new in edits.items():
|
|
if old not in text:
|
|
raise RuntimeError(f"Expected Resemble Enhance source not found in {path}: {old!r}")
|
|
text = text.replace(old, new, 1)
|
|
path.write_text(text, encoding="utf-8")
|