Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions .devcontainer/devDockerfile
Original file line number Diff line number Diff line change
@@ -0,0 +1,11 @@
ARG NEMO_VERSION=25.11

FROM nvcr.io/nvidia/nemo:${NEMO_VERSION}

LABEL \
maintainer="NeMoTTS team" \
description="NeMoTTS Dev Container"

WORKDIR /workspace/NeMo

USER root
26 changes: 26 additions & 0 deletions .devcontainer/devcontainer.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
{
"name": "NeMoTTS Dev Container",
"build": {
"dockerfile": "../.devcontainer/devDockerfile"
},
"runArgs": [
"--gpus",
"\"device=0,1\"",
"--ipc=host",
"--name",
"${localEnv:USER}_nemo_tts_devcontainer_easymagpietts_online_cfg_distillation"
],
"workspaceFolder": "/workspace/NeMo",
"mounts": [
"source=${localWorkspaceFolder},target=/workspace/NeMo,type=bind"
],
"remoteUser": "root",
"postCreateCommand": "pip install -e /workspace/NeMo --no-deps && pip install git+https://github.com/sarulab-speech/UTMOSv2.git@v1.2.1",
"customizations": {
"vscode": {
"settings": {
"terminal.integrated.shell.linux": "/bin/bash"
}
}
}
}
16 changes: 13 additions & 3 deletions examples/tts/easy_magpietts.py
100644 → 100755
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,17 @@
import torch.multiprocessing as mp
from omegaconf import OmegaConf, open_dict

from nemo.collections.tts.models import EasyMagpieTTSModel, EasyMagpieTTSModelOnlinePO
from nemo.collections.tts.models import EasyMagpieCFGDistillation, EasyMagpieTTSModel, EasyMagpieTTSModelOnlinePO
from nemo.core.config import hydra_runner
from nemo.utils import logging
from nemo.utils.exp_manager import exp_manager

_TRAIN_MODES: list[str] = [
"train",
"online_cfg_distillation_train",
"onlinepo_train",
]


@hydra_runner(config_path="conf/magpietts", config_name="easy_magpietts")
def main(cfg):
Expand All @@ -43,8 +49,12 @@ def main(cfg):
exp_manager(trainer, cfg.get("exp_manager", None))

mode = cfg.get('mode', 'train')
train_modes_msg = ", ".join(_TRAIN_MODES)

if mode == 'train':
model = EasyMagpieTTSModel(cfg=cfg.model, trainer=trainer)
elif mode == "online_cfg_distillation_train":
model = EasyMagpieCFGDistillation(cfg=cfg.model, trainer=trainer)
elif mode == 'onlinepo_train':
model_cfg = cfg.model
with open_dict(model_cfg):
Expand All @@ -53,11 +63,11 @@ def main(cfg):
elif mode == 'test':
model = EasyMagpieTTSModel(cfg=cfg.model, trainer=trainer)
else:
raise NotImplementedError(f"Only train, onlinepo_train and test modes are supported. Got {mode}")
raise NotImplementedError(f"Only {train_modes_msg} and test modes are supported. Got {mode}.")

model.maybe_init_from_pretrained_checkpoint(cfg=cfg)

if mode in ['train', 'onlinepo_train']:
if mode in _TRAIN_MODES:
trainer.fit(model)
elif mode == 'test':
trainer.test(model)
Expand Down
2 changes: 2 additions & 0 deletions nemo/collections/tts/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from nemo.collections.tts.models.aligner import AlignerModel
from nemo.collections.tts.models.audio_codec import AudioCodecModel
from nemo.collections.tts.models.easy_magpietts import EasyMagpieTTSModel
from nemo.collections.tts.models.easy_magpietts_cfg_distillation import EasyMagpieCFGDistillation
from nemo.collections.tts.models.easy_magpietts_inference import EasyMagpieTTSInferenceModel
from nemo.collections.tts.models.easy_magpietts_preference_optimization import EasyMagpieTTSModelOnlinePO
from nemo.collections.tts.models.fastpitch import FastPitchModel
Expand Down Expand Up @@ -45,4 +46,5 @@
"MagpieTTSModelOfflinePODataGen",
"MagpieTTSModelOfflinePO",
"MagpieTTSModelOnlinePO",
"EasyMagpieCFGDistillation",
]
Loading
Loading