first commit
This commit is contained in:
@@ -0,0 +1,166 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
SRC_ROOT = REPO_ROOT / "src"
|
||||
|
||||
for import_root in (REPO_ROOT, SRC_ROOT):
|
||||
import_root_str = str(import_root)
|
||||
if import_root_str not in sys.path:
|
||||
sys.path.insert(0, import_root_str)
|
||||
|
||||
import soundfile as sf # noqa: E402
|
||||
import torch # noqa: E402
|
||||
from loguru import logger # noqa: E402
|
||||
|
||||
from dots_tts.utils.logging import configure_logging # noqa: E402
|
||||
from dots_tts.runtime_double_streaming import ( # noqa: E402
|
||||
DotsTtsRuntimeDoubleStreaming,
|
||||
)
|
||||
from dots_tts.utils.text import normalize_text # noqa: E402
|
||||
from dots_tts.utils.util import seed_everything # noqa: E402
|
||||
|
||||
|
||||
def parse_args(argv=None):
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Temporary example for dots.tts double streaming session API."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model-name-or-path",
|
||||
required=True,
|
||||
help="Local pretrained directory or Hugging Face repo id",
|
||||
)
|
||||
parser.add_argument("--text", required=True, help="Input text")
|
||||
parser.add_argument("--output", default="double_streaming.wav", help="Output wav path")
|
||||
parser.add_argument(
|
||||
"--prompt-audio",
|
||||
default=None,
|
||||
help="Optional reference audio for ref_audio_only speaker conditioning",
|
||||
)
|
||||
parser.add_argument("--revision", default=None, help="Optional Hugging Face revision")
|
||||
parser.add_argument("--cache-dir", default=None, help="Optional Hugging Face cache dir")
|
||||
parser.add_argument("--precision", default="bfloat16", help="Inference precision")
|
||||
parser.add_argument(
|
||||
"--optimize",
|
||||
action="store_true",
|
||||
help="Enable inference optimization and warmup",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=42,
|
||||
help="Random seed.",
|
||||
)
|
||||
parser.add_argument("--ode-method", default="euler", help="ODE solver method")
|
||||
parser.add_argument("--num-steps", type=int, default=10, help="Diffusion sampling steps")
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=1.2,
|
||||
help="Classifier-free guidance scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--eos-threshold",
|
||||
type=float,
|
||||
default=0.8,
|
||||
help="EOS stop threshold for finish_text() tail decode",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-generate-length",
|
||||
type=int,
|
||||
default=500,
|
||||
help="Maximum number of decoded audio patches in double streaming",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--normalize-text",
|
||||
action="store_true",
|
||||
help="Normalize text before tokenizer encode",
|
||||
)
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def _prepare_text(text: str, *, normalize: bool) -> str:
|
||||
prepared = text.strip()
|
||||
if normalize:
|
||||
prepared = normalize_text(prepared)
|
||||
if not prepared:
|
||||
raise ValueError("Input text is empty after preprocessing.")
|
||||
return prepared
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
configure_logging()
|
||||
args = parse_args(argv)
|
||||
seed_everything(args.seed)
|
||||
|
||||
output_path = Path(args.output)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
runtime = DotsTtsRuntimeDoubleStreaming.from_pretrained(
|
||||
args.model_name_or_path,
|
||||
revision=args.revision,
|
||||
cache_dir=args.cache_dir,
|
||||
precision=args.precision,
|
||||
optimize=args.optimize,
|
||||
max_generate_length=args.max_generate_length,
|
||||
)
|
||||
prepared_text = _prepare_text(args.text, normalize=args.normalize_text)
|
||||
text_token_ids = runtime.model.tokenizer.encode(
|
||||
prepared_text,
|
||||
add_special_tokens=False,
|
||||
)
|
||||
if not text_token_ids:
|
||||
raise ValueError("Tokenizer produced no text tokens.")
|
||||
|
||||
logger.info(
|
||||
"Double streaming example started: text_len={} text_token_count={} output={}",
|
||||
len(prepared_text),
|
||||
len(text_token_ids),
|
||||
output_path,
|
||||
)
|
||||
|
||||
session = runtime.start_double_streaming(
|
||||
prompt_audio_path=args.prompt_audio,
|
||||
ode_method=args.ode_method,
|
||||
num_steps=args.num_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
eos_threshold=args.eos_threshold,
|
||||
)
|
||||
|
||||
chunks: list[torch.Tensor] = []
|
||||
for index, token_id in enumerate(text_token_ids, start=1):
|
||||
chunk = session.push_text_token(token_id)
|
||||
logger.info(
|
||||
"Double streaming step: token_index={} token_id={} emitted_audio={}",
|
||||
index,
|
||||
token_id,
|
||||
chunk is not None,
|
||||
)
|
||||
if chunk is not None:
|
||||
chunks.append(chunk.detach().cpu())
|
||||
|
||||
for chunk in session.finish_text():
|
||||
chunks.append(chunk.detach().cpu())
|
||||
|
||||
if not chunks:
|
||||
raise RuntimeError("Double streaming produced no audio chunks.")
|
||||
|
||||
audio = torch.cat(chunks, dim=-1)
|
||||
sf.write(
|
||||
output_path,
|
||||
audio.float().squeeze().numpy(),
|
||||
runtime.sample_rate,
|
||||
)
|
||||
logger.info(
|
||||
"Double streaming example completed: output={} chunk_count={} samples={}",
|
||||
output_path,
|
||||
len(chunks),
|
||||
audio.shape[-1],
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,134 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import json
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
REPO_ID = "alibabasglab/LJSpeech-1.1-48kHz"
|
||||
ARCHIVE_NAME = "LJSpeech-1.1-48kHz.tar.bz2"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--cache-dir",
|
||||
type=Path,
|
||||
default=REPO_ROOT / "downloaded_data" / "hf_cache",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--extract-dir",
|
||||
type=Path,
|
||||
default=REPO_ROOT / "downloaded_data" / "hf_cache",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-dir",
|
||||
type=Path,
|
||||
default=REPO_ROOT / "downloaded_data",
|
||||
)
|
||||
parser.add_argument("--valid-size", type=int, default=100)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.valid_size < 0:
|
||||
raise ValueError("valid_size must be >= 0")
|
||||
|
||||
cache_dir = args.cache_dir.resolve()
|
||||
extract_dir = args.extract_dir.resolve()
|
||||
output_dir = args.output_dir.resolve()
|
||||
train_manifest_path = output_dir / "ljspeech_48khz_manifest_train.jsonl"
|
||||
valid_manifest_path = output_dir / "ljspeech_48khz_manifest_valid.jsonl"
|
||||
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
extract_dir.mkdir(parents=True, exist_ok=True)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
archive_path = Path(
|
||||
hf_hub_download(
|
||||
repo_id=REPO_ID,
|
||||
repo_type="dataset",
|
||||
filename=ARCHIVE_NAME,
|
||||
local_dir=str(cache_dir),
|
||||
)
|
||||
)
|
||||
|
||||
dataset_root = extract_dir / "LJSpeech-1.1-48kHz"
|
||||
if not dataset_root.exists():
|
||||
print("extracting archive...")
|
||||
subprocess.run(
|
||||
[
|
||||
"tar",
|
||||
"-xjf",
|
||||
str(archive_path),
|
||||
"-C",
|
||||
str(extract_dir),
|
||||
"--checkpoint=2000",
|
||||
"--checkpoint-action=echo=extracting...",
|
||||
],
|
||||
check=True,
|
||||
)
|
||||
|
||||
metadata_path = dataset_root / "metadata.csv"
|
||||
audio_dir = dataset_root / "wavs" / "MossFormer2_SR_48K"
|
||||
|
||||
if not metadata_path.is_file():
|
||||
raise FileNotFoundError(f"metadata.csv not found: {metadata_path}")
|
||||
if not audio_dir.is_dir():
|
||||
raise FileNotFoundError(f"audio dir not found: {audio_dir}")
|
||||
|
||||
train_count = 0
|
||||
valid_count = 0
|
||||
with (
|
||||
metadata_path.open("r", encoding="utf-8", newline="") as fin,
|
||||
train_manifest_path.open("w", encoding="utf-8") as train_fout,
|
||||
valid_manifest_path.open("w", encoding="utf-8") as valid_fout,
|
||||
):
|
||||
reader = csv.reader(fin, delimiter="|")
|
||||
for row in reader:
|
||||
if not row:
|
||||
continue
|
||||
|
||||
fid = row[0].strip()
|
||||
text = (
|
||||
row[2].strip() if len(row) >= 3 and row[2].strip() else row[1].strip()
|
||||
)
|
||||
audio_path = (audio_dir / f"{fid}.wav").resolve()
|
||||
|
||||
if not audio_path.is_file():
|
||||
raise FileNotFoundError(f"audio not found: {audio_path}")
|
||||
|
||||
record = json.dumps(
|
||||
{
|
||||
"fid": fid,
|
||||
"audio": str(audio_path),
|
||||
"text": text,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
if valid_count < args.valid_size:
|
||||
valid_fout.write(record)
|
||||
valid_fout.write("\n")
|
||||
valid_count += 1
|
||||
else:
|
||||
train_fout.write(record)
|
||||
train_fout.write("\n")
|
||||
train_count += 1
|
||||
|
||||
print(f"archive: {archive_path}")
|
||||
print(f"dataset_root: {dataset_root}")
|
||||
print(f"train_manifest: {train_manifest_path}")
|
||||
print(f"valid_manifest: {valid_manifest_path}")
|
||||
print(f"train_records: {train_count}")
|
||||
print(f"valid_records: {valid_count}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+773
@@ -0,0 +1,773 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import math
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import yaml
|
||||
from accelerate import Accelerator
|
||||
from accelerate.utils import DistributedDataParallelKwargs, ProjectConfiguration
|
||||
from torch.optim import AdamW
|
||||
from transformers import get_cosine_schedule_with_warmup
|
||||
|
||||
from dots_tts.config import app as app_config
|
||||
from dots_tts.data import builders as data_module
|
||||
from dots_tts.models.dots_tts import model as dots_tts_model
|
||||
from dots_tts.training import checkpoint as train_checkpoint
|
||||
from dots_tts.training import losses as loss_ops
|
||||
from dots_tts.training import utils as train_utils
|
||||
from dots_tts.utils import util as util_module
|
||||
|
||||
_EMPTY_EPOCH_TOLERANCE = 32
|
||||
_DEBUG_BATCH_LIMIT = 3
|
||||
_DEBUG_GRAD_EARLY_STEP_LIMIT = 3
|
||||
|
||||
|
||||
# region Training Step State
|
||||
@dataclass(slots=True)
|
||||
class _PreparedTrainingStep:
|
||||
micro_batches: list[dict]
|
||||
consumed_counts: list[int]
|
||||
global_denominators: dict[str, float]
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _AccumulatedTrainingStep:
|
||||
loss_totals: dict[str, float]
|
||||
loss_denominators: dict[str, float]
|
||||
source_loss_totals: dict[str, dict[str, float]]
|
||||
source_loss_denominators: dict[str, dict[str, float]]
|
||||
completed_optimizer_step: bool
|
||||
grad_norm: torch.Tensor | None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _CompletedTrainingStep:
|
||||
reduced_metrics: dict[str, float]
|
||||
learning_rate: float
|
||||
grad_norm_value: float
|
||||
|
||||
# endregion Training Step State
|
||||
|
||||
class DotsTtsTrainingRun:
|
||||
# region Lifecycle
|
||||
def __init__(self, cfg: app_config.AppConfig, *, debug_enabled: bool = False):
|
||||
self.cfg = cfg
|
||||
self.progress = train_utils.TrainProgress()
|
||||
self.max_train_steps = int(cfg.train.max_train_steps)
|
||||
self.grad_accumulation_steps = int(cfg.train.gradient_accumulation_steps)
|
||||
self.last_validation_step: int | None = None
|
||||
self.consecutive_empty_epochs = 0
|
||||
self.saved_latest_checkpoint = False
|
||||
self._last_log_step = 0
|
||||
self._last_log_time = 0.0
|
||||
self._debug_enabled = bool(debug_enabled)
|
||||
self._debug_batch_count = 0
|
||||
self._debug_audio_sample_rate = int(self.cfg.train_data.train_audio_sample_rate)
|
||||
|
||||
project_config = ProjectConfiguration(
|
||||
project_dir=self.cfg.train.output_dir,
|
||||
total_limit=self.cfg.train.max_checkpoints_to_keep,
|
||||
)
|
||||
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=False)
|
||||
self.accelerator = Accelerator(
|
||||
kwargs_handlers=[ddp_kwargs],
|
||||
gradient_accumulation_steps=self.grad_accumulation_steps,
|
||||
log_with="tensorboard",
|
||||
project_config=project_config,
|
||||
step_scheduler_with_optimizer=False,
|
||||
)
|
||||
|
||||
util_module.seed_everything(self.cfg.train.seed)
|
||||
|
||||
model = dots_tts_model.DotsTtsModel.from_pretrained(
|
||||
self.cfg.train.pretrained_model_path
|
||||
)
|
||||
# model.set_cfg_droprate(
|
||||
# cfg_droprate=self.cfg.train.cfg_droprate,
|
||||
# xvec_drop_rate=self.cfg.train.xvec_drop_rate,
|
||||
# )
|
||||
optimizer = AdamW(
|
||||
(param for param in model.parameters() if param.requires_grad),
|
||||
lr=self.cfg.train.learning_rate,
|
||||
weight_decay=self.cfg.train.weight_decay,
|
||||
)
|
||||
scheduler = get_cosine_schedule_with_warmup(
|
||||
optimizer,
|
||||
num_warmup_steps=self.cfg.train.warmup_steps,
|
||||
num_training_steps=self.max_train_steps,
|
||||
)
|
||||
self.model, self.optimizer, self.scheduler = self.accelerator.prepare(
|
||||
model,
|
||||
optimizer,
|
||||
scheduler,
|
||||
)
|
||||
self.unwrapped_model = self.accelerator.unwrap_model(self.model)
|
||||
expected_sample_rate = int(self.unwrapped_model.config.vocoder.sample_rate)
|
||||
expected_audio_samples_per_llm_token = (
|
||||
int(self.unwrapped_model.hop_size) * int(self.unwrapped_model.config.patch_size)
|
||||
)
|
||||
if int(self.cfg.train_data.train_audio_sample_rate) != expected_sample_rate:
|
||||
raise ValueError(
|
||||
f"train_data.train_audio_sample_rate={int(self.cfg.train_data.train_audio_sample_rate)} "
|
||||
f"does not match the pretrained model sample rate {expected_sample_rate}."
|
||||
)
|
||||
if (
|
||||
int(self.cfg.train_data.audio_samples_per_llm_token)
|
||||
!= expected_audio_samples_per_llm_token
|
||||
):
|
||||
raise ValueError(
|
||||
"train_data.audio_samples_per_llm_token="
|
||||
f"{int(self.cfg.train_data.audio_samples_per_llm_token)} "
|
||||
"does not match the pretrained model audio token contract "
|
||||
f"{expected_audio_samples_per_llm_token}."
|
||||
)
|
||||
if self.cfg.val_data is not None:
|
||||
if int(self.cfg.val_data.train_audio_sample_rate) != expected_sample_rate:
|
||||
raise ValueError(
|
||||
f"val_data.train_audio_sample_rate={int(self.cfg.val_data.train_audio_sample_rate)} "
|
||||
f"does not match the pretrained model sample rate {expected_sample_rate}."
|
||||
)
|
||||
if (
|
||||
int(self.cfg.val_data.audio_samples_per_llm_token)
|
||||
!= expected_audio_samples_per_llm_token
|
||||
):
|
||||
raise ValueError(
|
||||
"val_data.audio_samples_per_llm_token="
|
||||
f"{int(self.cfg.val_data.audio_samples_per_llm_token)} "
|
||||
"does not match the pretrained model audio token contract "
|
||||
f"{expected_audio_samples_per_llm_token}."
|
||||
)
|
||||
|
||||
if self.accelerator.is_main_process:
|
||||
total_params = sum(param.numel() for param in self.unwrapped_model.parameters())
|
||||
trainable_params = sum(
|
||||
param.numel()
|
||||
for param in self.unwrapped_model.parameters()
|
||||
if param.requires_grad
|
||||
)
|
||||
self.accelerator.print(f"Total parameters: {total_params:,}")
|
||||
self.accelerator.print(f"Trainable parameters: {trainable_params:,}")
|
||||
self.accelerator.print(
|
||||
f"Distributed type: {self.accelerator.distributed_type}"
|
||||
)
|
||||
|
||||
tokenizer = self.unwrapped_model.tokenizer
|
||||
self.tokenizer = tokenizer
|
||||
train_dataset = data_module.build_training_dataset(
|
||||
self.cfg.train_data,
|
||||
tokenizer=tokenizer,
|
||||
seed=int(self.cfg.train.seed),
|
||||
accelerator=self.accelerator,
|
||||
)
|
||||
self.train_loader = data_module.build_training_dataloader(
|
||||
train_dataset,
|
||||
self.cfg.train_data,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
|
||||
self.val_loader = None
|
||||
if (
|
||||
self.cfg.train.eval_interval is not None
|
||||
or self.cfg.train.run_eval_on_start
|
||||
):
|
||||
if self.cfg.val_data is None:
|
||||
raise ValueError(
|
||||
"Validation requires val_data when eval_interval or "
|
||||
"run_eval_on_start is enabled."
|
||||
)
|
||||
validation_data_cfg = self.cfg.val_data.model_copy(deep=True)
|
||||
validation_data_cfg.num_tokens_per_epoch = None
|
||||
val_dataset = data_module.build_validation_dataset(
|
||||
validation_data_cfg,
|
||||
tokenizer=tokenizer,
|
||||
seed=int(self.cfg.train.seed),
|
||||
accelerator=self.accelerator,
|
||||
)
|
||||
self.val_loader = data_module.build_validation_dataloader(
|
||||
val_dataset,
|
||||
validation_data_cfg,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
|
||||
self._resume_if_available()
|
||||
self.train_loader.set_epoch(self.progress.epoch)
|
||||
|
||||
def run(self) -> int:
|
||||
self.accelerator.init_trackers("dots_tts")
|
||||
self._write_run_config()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
try:
|
||||
if self.cfg.train.run_eval_on_start:
|
||||
self._run_validation()
|
||||
self.last_validation_step = self.progress.global_step
|
||||
|
||||
self._last_log_step = self.progress.global_step
|
||||
self._last_log_time = time.perf_counter()
|
||||
|
||||
while self.progress.global_step < self.max_train_steps:
|
||||
self._run_training_step()
|
||||
|
||||
if (
|
||||
self.cfg.train.eval_interval is not None
|
||||
and self.val_loader is not None
|
||||
and self.progress.global_step > 0
|
||||
and self.last_validation_step != self.progress.global_step
|
||||
):
|
||||
self._run_validation()
|
||||
|
||||
if not self.saved_latest_checkpoint:
|
||||
self._save_checkpoint(float(self.optimizer.param_groups[0]["lr"]))
|
||||
return 0
|
||||
finally:
|
||||
try:
|
||||
self._close_data_streams()
|
||||
finally:
|
||||
self.accelerator.end_training()
|
||||
|
||||
def _write_run_config(self) -> None:
|
||||
if not bool(getattr(self.accelerator, "is_main_process", True)):
|
||||
return
|
||||
output_dir = Path(self.cfg.train.output_dir)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
config_path = output_dir / "config.yml"
|
||||
with config_path.open("w", encoding="utf-8") as fout:
|
||||
yaml.safe_dump(
|
||||
self.cfg.to_dict(),
|
||||
fout,
|
||||
sort_keys=False,
|
||||
allow_unicode=True,
|
||||
)
|
||||
|
||||
def _close_data_streams(self) -> None:
|
||||
for loader_name in ("train_loader", "val_loader"):
|
||||
loader = getattr(self, loader_name, None)
|
||||
close = getattr(loader, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
setattr(self, loader_name, None)
|
||||
|
||||
def _resume_if_available(self) -> None:
|
||||
try:
|
||||
resume_dir = train_checkpoint.resolve_latest_train_checkpoint(
|
||||
self.cfg.train.output_dir
|
||||
)
|
||||
except FileNotFoundError:
|
||||
return
|
||||
|
||||
resume_state = train_checkpoint.load_train_checkpoint(
|
||||
self.accelerator,
|
||||
self.model,
|
||||
self.optimizer,
|
||||
self.progress,
|
||||
resume_dir,
|
||||
self.scheduler,
|
||||
)
|
||||
saved_max_train_steps = int(resume_state["scheduler_state"]["max_train_steps"])
|
||||
if saved_max_train_steps != self.max_train_steps:
|
||||
self.accelerator.print(
|
||||
"Warning: resumed scheduler was saved with "
|
||||
f"max_train_steps={saved_max_train_steps}, but current run uses "
|
||||
f"{self.max_train_steps}."
|
||||
)
|
||||
|
||||
self.train_loader.load_state_dict(resume_state["data_state"])
|
||||
self.accelerator.print(
|
||||
"Resumed training from "
|
||||
f"{resume_dir} at step {self.progress.global_step}. "
|
||||
"Restored committed data state. "
|
||||
"In-memory prefetch and batching state is rebuilt on restart, so only "
|
||||
"committed sample progress is resumed."
|
||||
)
|
||||
# endregion Lifecycle
|
||||
|
||||
# region Training Step Pipeline
|
||||
def _run_training_step(self) -> None:
|
||||
try:
|
||||
self.model.train()
|
||||
# Stage 1: collect one synchronized accumulation window and its
|
||||
# normalization factors before touching model state.
|
||||
prepared_step = self._prepare_training_step()
|
||||
|
||||
# Stage 2: run forward/backward over the prepared micro-batches and
|
||||
# accumulate overall/source statistics for the completed optimizer step.
|
||||
accumulated_step = self._accumulate_training_step(prepared_step)
|
||||
|
||||
# Stage 3: advance counters, reduce metrics, then trigger side effects
|
||||
# (logging, validation, checkpointing) only after a real optimizer step.
|
||||
self._apply_consumed_counts(prepared_step.consumed_counts)
|
||||
if not accumulated_step.completed_optimizer_step:
|
||||
return
|
||||
completed_step = self._finalize_completed_training_step(accumulated_step)
|
||||
if train_utils.should_log_training_step(
|
||||
self.progress.global_step,
|
||||
int(self.cfg.train.log_interval),
|
||||
):
|
||||
reduced_by_source = train_utils.reduce_source_metrics(
|
||||
accumulated_step.source_loss_totals,
|
||||
accumulated_step.source_loss_denominators,
|
||||
device=self.accelerator.device,
|
||||
loss_config=self.cfg.loss,
|
||||
)
|
||||
current_time = time.perf_counter()
|
||||
report = train_utils.build_train_step_report(
|
||||
completed_step.reduced_metrics,
|
||||
learning_rate=completed_step.learning_rate,
|
||||
grad_norm=completed_step.grad_norm_value,
|
||||
current_time=current_time,
|
||||
last_log_step=self._last_log_step,
|
||||
last_log_time=self._last_log_time,
|
||||
progress=self.progress,
|
||||
max_train_steps=self.max_train_steps,
|
||||
reduced_by_source=reduced_by_source,
|
||||
)
|
||||
self.accelerator.log(
|
||||
report.log_values,
|
||||
step=self.progress.global_step,
|
||||
)
|
||||
self.accelerator.print(report.console_line)
|
||||
self._last_log_step = self.progress.global_step
|
||||
self._last_log_time = current_time
|
||||
|
||||
if (
|
||||
self.cfg.train.eval_interval is not None
|
||||
and self.progress.global_step % self.cfg.train.eval_interval == 0
|
||||
):
|
||||
self._run_validation()
|
||||
self.last_validation_step = self.progress.global_step
|
||||
|
||||
if self.progress.global_step % self.cfg.train.save_interval == 0:
|
||||
self._save_checkpoint(completed_step.learning_rate)
|
||||
self.saved_latest_checkpoint = True
|
||||
except BaseException as exc:
|
||||
train_utils.abort_on_out_of_memory(
|
||||
exc,
|
||||
stage="train",
|
||||
batch=None,
|
||||
progress=self.progress,
|
||||
device=self.accelerator.device,
|
||||
process_index=int(getattr(self.accelerator, "process_index", 0)),
|
||||
num_processes=int(getattr(self.accelerator, "num_processes", 1)),
|
||||
)
|
||||
raise
|
||||
|
||||
def _prepare_training_step(self) -> _PreparedTrainingStep:
|
||||
micro_batches: list[dict] = []
|
||||
local_denominators: dict[str, float] = {}
|
||||
|
||||
while len(micro_batches) < self.grad_accumulation_steps:
|
||||
batch, has_batch = self.train_loader.peek_batch()
|
||||
if train_utils.any_rank_true(not has_batch, device=self.accelerator.device):
|
||||
self._advance_epoch_after_empty_batch(has_local_batch=has_batch)
|
||||
continue
|
||||
|
||||
self.consecutive_empty_epochs = 0
|
||||
self.train_loader.commit_batch()
|
||||
prepared_batch = self.unwrapped_model.prepare_training_batch(batch)
|
||||
self._maybe_debug_training_batch(prepared_batch)
|
||||
batch_denominators = loss_ops.to_host_named_scalars(
|
||||
loss_ops.collapse_loss_masks(prepared_batch["loss_masks"])
|
||||
)
|
||||
if not local_denominators:
|
||||
local_denominators = {name: 0.0 for name in batch_denominators}
|
||||
loss_ops.accumulate_named_scalars_(local_denominators, batch_denominators)
|
||||
micro_batches.append(prepared_batch)
|
||||
|
||||
consumed_counts = train_utils.sum_integer_counters_across_ranks(
|
||||
[
|
||||
sum(int(batch["input_ids_lengths"].sum().item()) for batch in micro_batches),
|
||||
sum(int(batch["num_audio_tokens"].sum().item()) for batch in micro_batches),
|
||||
sum(int(batch["num_text_tokens"].sum().item()) for batch in micro_batches),
|
||||
],
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
global_denominators = loss_ops.sum_named_scalars_across_ranks(
|
||||
local_denominators,
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
return _PreparedTrainingStep(
|
||||
micro_batches=micro_batches,
|
||||
consumed_counts=consumed_counts,
|
||||
global_denominators=global_denominators,
|
||||
)
|
||||
|
||||
def _advance_epoch_after_empty_batch(self, *, has_local_batch: bool) -> None:
|
||||
if has_local_batch:
|
||||
self.train_loader.discard_batch()
|
||||
self.progress.epoch += 1
|
||||
self.train_loader.set_epoch(self.progress.epoch)
|
||||
self.consecutive_empty_epochs += 1
|
||||
if self.consecutive_empty_epochs > _EMPTY_EPOCH_TOLERANCE:
|
||||
raise RuntimeError(
|
||||
"Unable to obtain a synchronized training batch across ranks. "
|
||||
"Check shard assignment, dataset size, and filtering constraints."
|
||||
)
|
||||
|
||||
def _accumulate_training_step(
|
||||
self,
|
||||
prepared_step: _PreparedTrainingStep,
|
||||
) -> _AccumulatedTrainingStep:
|
||||
accumulated_loss_totals: dict[str, float] = {}
|
||||
accumulated_loss_denominators: dict[str, float] = {}
|
||||
accumulated_source_loss_totals: dict[str, dict[str, float]] = {}
|
||||
accumulated_source_loss_denominators: dict[str, dict[str, float]] = {}
|
||||
completed_optimizer_step = False
|
||||
grad_norm = None
|
||||
|
||||
for batch in prepared_step.micro_batches:
|
||||
batch = train_utils.move_to_device(batch, self.accelerator.device)
|
||||
with self.accelerator.accumulate(self.model):
|
||||
with self.accelerator.autocast():
|
||||
loss_terms = self.model(batch)
|
||||
loss = loss_ops.compute_gradient_loss(
|
||||
loss_terms,
|
||||
global_normalizers=prepared_step.global_denominators,
|
||||
loss_config=self.cfg.loss,
|
||||
ddp_world_size=int(self.accelerator.num_processes),
|
||||
gradient_accumulation_steps=self.grad_accumulation_steps,
|
||||
)
|
||||
|
||||
batch_loss_totals, batch_loss_denominators = (
|
||||
loss_ops.collapse_loss_terms(loss_terms)
|
||||
)
|
||||
batch_loss_totals = loss_ops.to_host_named_scalars(batch_loss_totals)
|
||||
batch_loss_denominators = loss_ops.to_host_named_scalars(
|
||||
batch_loss_denominators
|
||||
)
|
||||
if not accumulated_loss_totals:
|
||||
accumulated_loss_totals = {name: 0.0 for name in batch_loss_totals}
|
||||
accumulated_loss_denominators = {
|
||||
name: 0.0 for name in batch_loss_denominators
|
||||
}
|
||||
loss_ops.accumulate_named_scalars_(
|
||||
accumulated_loss_totals,
|
||||
batch_loss_totals,
|
||||
)
|
||||
loss_ops.accumulate_named_scalars_(
|
||||
accumulated_loss_denominators,
|
||||
batch_loss_denominators,
|
||||
)
|
||||
|
||||
batch_source_totals, batch_source_denominators = (
|
||||
loss_ops.collapse_loss_terms_by_source(
|
||||
loss_terms,
|
||||
source_names=batch["source_names"],
|
||||
)
|
||||
)
|
||||
loss_ops.accumulate_grouped_named_scalars_(
|
||||
accumulated_source_loss_totals,
|
||||
batch_source_totals,
|
||||
)
|
||||
loss_ops.accumulate_grouped_named_scalars_(
|
||||
accumulated_source_loss_denominators,
|
||||
batch_source_denominators,
|
||||
)
|
||||
|
||||
self.accelerator.backward(loss)
|
||||
if self.accelerator.sync_gradients:
|
||||
grad_norm = self.accelerator.clip_grad_norm_(
|
||||
self.model.parameters(),
|
||||
self.cfg.train.grad_clip_norm,
|
||||
)
|
||||
self._maybe_print_gradient_debug(grad_norm)
|
||||
self.optimizer.step()
|
||||
completed_optimizer_step = (
|
||||
not self.accelerator.optimizer_step_was_skipped
|
||||
)
|
||||
if completed_optimizer_step:
|
||||
self.scheduler.step()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
batch.clear()
|
||||
|
||||
return _AccumulatedTrainingStep(
|
||||
loss_totals=accumulated_loss_totals,
|
||||
loss_denominators=accumulated_loss_denominators,
|
||||
source_loss_totals=accumulated_source_loss_totals,
|
||||
source_loss_denominators=accumulated_source_loss_denominators,
|
||||
completed_optimizer_step=completed_optimizer_step,
|
||||
grad_norm=grad_norm,
|
||||
)
|
||||
|
||||
def _apply_consumed_counts(self, consumed_counts: list[int]) -> None:
|
||||
self.progress.total_tokens += consumed_counts[0]
|
||||
self.progress.audio_tokens += consumed_counts[1]
|
||||
self.progress.text_tokens += consumed_counts[2]
|
||||
|
||||
def _finalize_completed_training_step(
|
||||
self,
|
||||
accumulated_step: _AccumulatedTrainingStep,
|
||||
) -> _CompletedTrainingStep:
|
||||
if not accumulated_step.loss_totals or not accumulated_step.loss_denominators:
|
||||
raise RuntimeError("Training step produced no accumulated loss totals.")
|
||||
if all(
|
||||
float(value) == 0.0 for value in accumulated_step.loss_denominators.values()
|
||||
):
|
||||
raise RuntimeError("Accumulated training step produced no loss statistics.")
|
||||
|
||||
self.progress.global_step += 1
|
||||
self.saved_latest_checkpoint = False
|
||||
|
||||
reduced_totals = loss_ops.sum_named_scalars_across_ranks(
|
||||
accumulated_step.loss_totals,
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
reduced_denominators = loss_ops.sum_named_scalars_across_ranks(
|
||||
accumulated_step.loss_denominators,
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
reduced_metrics = loss_ops.reduce_loss_statistics(
|
||||
reduced_totals,
|
||||
reduced_denominators,
|
||||
loss_config=self.cfg.loss,
|
||||
)
|
||||
learning_rate = float(self.optimizer.param_groups[0]["lr"])
|
||||
grad_norm_value = (
|
||||
math.nan
|
||||
if accumulated_step.grad_norm is None
|
||||
else float(accumulated_step.grad_norm.detach().float().item())
|
||||
)
|
||||
return _CompletedTrainingStep(
|
||||
reduced_metrics=reduced_metrics,
|
||||
learning_rate=learning_rate,
|
||||
grad_norm_value=grad_norm_value,
|
||||
)
|
||||
# endregion Training Step Pipeline
|
||||
|
||||
# region Validation
|
||||
def _run_validation(self) -> None:
|
||||
try:
|
||||
if self.val_loader is None:
|
||||
raise ValueError(
|
||||
"Validation requested, but validation loader was not initialized."
|
||||
)
|
||||
self.val_loader.set_epoch(0)
|
||||
|
||||
was_training = bool(self.model.training)
|
||||
self.model.eval()
|
||||
|
||||
overall_loss_totals = None
|
||||
overall_loss_denominators = None
|
||||
source_loss_totals: dict[str, dict[str, float]] = {}
|
||||
source_loss_denominators: dict[str, dict[str, float]] = {}
|
||||
processed_batches = 0
|
||||
|
||||
# Collect rank-local partial sums using the same batch preparation and
|
||||
# loss aggregation path as training.
|
||||
with torch.no_grad():
|
||||
for batch_idx, batch in enumerate(self.val_loader):
|
||||
if (
|
||||
self.cfg.train.max_eval_batches is not None
|
||||
and batch_idx >= self.cfg.train.max_eval_batches
|
||||
):
|
||||
break
|
||||
|
||||
batch = self.unwrapped_model.prepare_training_batch(batch)
|
||||
batch = train_utils.move_to_device(batch, self.accelerator.device)
|
||||
|
||||
with self.accelerator.autocast():
|
||||
loss_terms = self.model(batch)
|
||||
|
||||
batch_loss_totals, batch_loss_denominators = (
|
||||
loss_ops.collapse_loss_terms(loss_terms)
|
||||
)
|
||||
batch_loss_totals = loss_ops.to_host_named_scalars(batch_loss_totals)
|
||||
batch_loss_denominators = loss_ops.to_host_named_scalars(
|
||||
batch_loss_denominators
|
||||
)
|
||||
if overall_loss_totals is None:
|
||||
overall_loss_totals = {name: 0.0 for name in batch_loss_totals}
|
||||
overall_loss_denominators = {
|
||||
name: 0.0 for name in batch_loss_denominators
|
||||
}
|
||||
loss_ops.accumulate_named_scalars_(
|
||||
overall_loss_totals,
|
||||
batch_loss_totals,
|
||||
)
|
||||
loss_ops.accumulate_named_scalars_(
|
||||
overall_loss_denominators,
|
||||
batch_loss_denominators,
|
||||
)
|
||||
|
||||
batch_source_totals, batch_source_denominators = (
|
||||
loss_ops.collapse_loss_terms_by_source(
|
||||
loss_terms,
|
||||
source_names=batch["source_names"],
|
||||
)
|
||||
)
|
||||
loss_ops.accumulate_grouped_named_scalars_(
|
||||
source_loss_totals,
|
||||
batch_source_totals,
|
||||
)
|
||||
loss_ops.accumulate_grouped_named_scalars_(
|
||||
source_loss_denominators,
|
||||
batch_source_denominators,
|
||||
)
|
||||
processed_batches += 1
|
||||
|
||||
# Merge rank-local partial sums with tensor reductions only. Validation
|
||||
# runs close to the training memory ceiling, so object collectives are
|
||||
# not acceptable here because NCCL materializes pickled payloads on GPU.
|
||||
processed_batches = train_utils.sum_integer_counters_across_ranks(
|
||||
[processed_batches],
|
||||
device=self.accelerator.device,
|
||||
)[0]
|
||||
overall_loss_totals = loss_ops.sum_named_scalars_across_ranks(
|
||||
overall_loss_totals or {},
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
overall_loss_denominators = loss_ops.sum_named_scalars_across_ranks(
|
||||
overall_loss_denominators or {},
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
source_loss_totals = loss_ops.sum_grouped_named_scalars_across_ranks(
|
||||
source_loss_totals,
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
source_loss_denominators = (
|
||||
loss_ops.sum_grouped_named_scalars_across_ranks(
|
||||
source_loss_denominators,
|
||||
device=self.accelerator.device,
|
||||
)
|
||||
)
|
||||
|
||||
if processed_batches <= 0:
|
||||
raise RuntimeError(
|
||||
"Validation produced no batches. Check validation data configuration."
|
||||
)
|
||||
if not overall_loss_totals or not overall_loss_denominators:
|
||||
raise RuntimeError("Validation produced no aggregate loss totals.")
|
||||
|
||||
reduced_metrics = loss_ops.reduce_loss_statistics(
|
||||
overall_loss_totals,
|
||||
overall_loss_denominators,
|
||||
loss_config=self.cfg.loss,
|
||||
)
|
||||
reduced_by_source = loss_ops.reduce_loss_statistics_by_source(
|
||||
source_loss_totals,
|
||||
source_loss_denominators,
|
||||
loss_config=self.cfg.loss,
|
||||
)
|
||||
|
||||
if was_training:
|
||||
self.model.train()
|
||||
|
||||
self.accelerator.log(
|
||||
train_utils.build_validation_log_dict(
|
||||
reduced_metrics,
|
||||
reduced_by_source=reduced_by_source,
|
||||
),
|
||||
step=self.progress.global_step,
|
||||
)
|
||||
self.accelerator.print(
|
||||
train_utils.format_validation_line(
|
||||
reduced_metrics,
|
||||
global_step=self.progress.global_step,
|
||||
reduced_by_source=reduced_by_source,
|
||||
)
|
||||
)
|
||||
except BaseException as exc:
|
||||
train_utils.abort_on_out_of_memory(
|
||||
exc,
|
||||
stage="validation",
|
||||
batch=None,
|
||||
progress=self.progress,
|
||||
device=self.accelerator.device,
|
||||
process_index=int(getattr(self.accelerator, "process_index", 0)),
|
||||
num_processes=int(getattr(self.accelerator, "num_processes", 1)),
|
||||
)
|
||||
raise
|
||||
# endregion Validation
|
||||
|
||||
# region Checkpointing
|
||||
def _save_checkpoint(self, learning_rate: float) -> None:
|
||||
train_checkpoint.save_train_checkpoint(
|
||||
self.accelerator,
|
||||
self.model,
|
||||
self.optimizer,
|
||||
self.progress,
|
||||
self.cfg.train.output_dir,
|
||||
self.cfg.train.max_checkpoints_to_keep,
|
||||
self.train_loader.state_dict(),
|
||||
{
|
||||
"type": "transformers_cosine_with_warmup",
|
||||
"global_step": int(self.progress.global_step),
|
||||
"base_lr": float(self.cfg.train.learning_rate),
|
||||
"current_lr": float(learning_rate),
|
||||
"warmup_steps": int(self.cfg.train.warmup_steps),
|
||||
"max_train_steps": int(self.max_train_steps),
|
||||
"state_dict": self.scheduler.state_dict(),
|
||||
},
|
||||
)
|
||||
# endregion Checkpointing
|
||||
|
||||
# region Debug Logging
|
||||
def _maybe_debug_training_batch(self, batch: dict[str, object]) -> None:
|
||||
if not bool(getattr(self, "_debug_enabled", False)):
|
||||
return
|
||||
if not bool(getattr(self.accelerator, "is_main_process", True)):
|
||||
return
|
||||
if self._debug_batch_count >= _DEBUG_BATCH_LIMIT:
|
||||
return
|
||||
|
||||
batch_index = self._debug_batch_count
|
||||
self._debug_batch_count += 1
|
||||
for line in train_utils.build_data_debug_lines(
|
||||
batch,
|
||||
batch_index=batch_index,
|
||||
tokenizer=self.tokenizer,
|
||||
sample_rate=self._debug_audio_sample_rate,
|
||||
):
|
||||
self.accelerator.print(line)
|
||||
|
||||
def _maybe_print_gradient_debug(self, grad_norm: torch.Tensor | None) -> None:
|
||||
if grad_norm is None:
|
||||
return
|
||||
if not train_utils.should_print_gradient_debug(
|
||||
debug_enabled=bool(getattr(self, "_debug_enabled", False)),
|
||||
is_main_process=bool(getattr(self.accelerator, "is_main_process", True)),
|
||||
next_global_step=self.progress.global_step + 1,
|
||||
log_interval=int(self.cfg.train.log_interval),
|
||||
early_step_limit=_DEBUG_GRAD_EARLY_STEP_LIMIT,
|
||||
):
|
||||
return
|
||||
for line in train_utils.build_gradient_debug_lines(
|
||||
self.unwrapped_model,
|
||||
global_step=self.progress.global_step + 1,
|
||||
grad_norm=float(grad_norm.detach().float().item()),
|
||||
grad_clip_norm=float(self.cfg.train.grad_clip_norm),
|
||||
):
|
||||
self.accelerator.print(line)
|
||||
# endregion Debug Logging
|
||||
|
||||
|
||||
# region CLI
|
||||
def parse_args(argv=None):
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Accelerate training entrypoint for dots.tts."
|
||||
)
|
||||
parser.add_argument("--config", default=app_config.DEFAULT_CONFIG_PATH)
|
||||
parser.add_argument(
|
||||
"--debug",
|
||||
action="store_true",
|
||||
help="Print training debug information.",
|
||||
)
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
args = parse_args(argv)
|
||||
return DotsTtsTrainingRun(
|
||||
app_config.load_config(args.config),
|
||||
debug_enabled=args.debug,
|
||||
).run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
# endregion CLI
|
||||
@@ -0,0 +1,956 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import yaml
|
||||
from accelerate import Accelerator
|
||||
from accelerate.utils import DistributedDataParallelKwargs, ProjectConfiguration
|
||||
from einops import rearrange
|
||||
from torch.optim import AdamW
|
||||
from train_dots_tts import DotsTtsTrainingRun
|
||||
from transformers import get_cosine_schedule_with_warmup
|
||||
|
||||
from dots_tts.config import app as app_config
|
||||
from dots_tts.data import builders as data_module
|
||||
from dots_tts.models.dots_tts import model as dots_tts_model
|
||||
from dots_tts.models.dots_tts.config import MeanFlowConfig
|
||||
from dots_tts.models.dots_tts.core import DotsTtsForwardOutput
|
||||
from dots_tts.modules.backbone.dit import DiT
|
||||
from dots_tts.training import checkpoint as train_checkpoint
|
||||
from dots_tts.training import utils as train_utils
|
||||
from dots_tts.utils import util as util_module
|
||||
|
||||
_ALLOWED_TEACHER_SOLVERS = ("euler", "midpoint", "rk4")
|
||||
_ALLOWED_CFG_DISTILL_MODES = ("natural", "fused")
|
||||
_ALLOWED_ANCHOR_TARGETS = ("formula", "teacher")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MeanFlowSettings:
|
||||
teacher_model_path: str | None
|
||||
teacher_steps: int = 8
|
||||
teacher_solver: str = "euler"
|
||||
cfg_distill_mode: str = "fused"
|
||||
distill_cfg_scale: float = 1.2
|
||||
anchor_prob: float = 0.5
|
||||
anchor_target: str = "formula"
|
||||
time_sampling_mean: float = -0.4
|
||||
time_sampling_std: float = 1.0
|
||||
train_all_parameters: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if int(self.teacher_steps) <= 0:
|
||||
raise ValueError("teacher_steps must be positive.")
|
||||
if self.teacher_solver not in _ALLOWED_TEACHER_SOLVERS:
|
||||
raise ValueError(
|
||||
f"teacher_solver must be one of {_ALLOWED_TEACHER_SOLVERS}, "
|
||||
f"got {self.teacher_solver!r}."
|
||||
)
|
||||
if self.cfg_distill_mode not in _ALLOWED_CFG_DISTILL_MODES:
|
||||
raise ValueError(
|
||||
"cfg_distill_mode must be one of "
|
||||
f"{_ALLOWED_CFG_DISTILL_MODES}, got {self.cfg_distill_mode!r}."
|
||||
)
|
||||
if self.anchor_target not in _ALLOWED_ANCHOR_TARGETS:
|
||||
raise ValueError(
|
||||
f"anchor_target must be one of {_ALLOWED_ANCHOR_TARGETS}, "
|
||||
f"got {self.anchor_target!r}."
|
||||
)
|
||||
if not 0.0 <= float(self.anchor_prob) <= 1.0:
|
||||
raise ValueError("anchor_prob must be in [0, 1].")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"teacher_model_path": self.teacher_model_path,
|
||||
"teacher_steps": int(self.teacher_steps),
|
||||
"teacher_solver": self.teacher_solver,
|
||||
"cfg_distill_mode": self.cfg_distill_mode,
|
||||
"distill_cfg_scale": float(self.distill_cfg_scale),
|
||||
"anchor_prob": float(self.anchor_prob),
|
||||
"anchor_target": self.anchor_target,
|
||||
"time_sampling_mean": float(self.time_sampling_mean),
|
||||
"time_sampling_std": float(self.time_sampling_std),
|
||||
"train_all_parameters": bool(self.train_all_parameters),
|
||||
}
|
||||
|
||||
|
||||
def enable_meanflow_student(model: dots_tts_model.DotsTtsModel) -> None:
|
||||
meanflow_config = MeanFlowConfig(enabled=True, use_duration_embedding=True)
|
||||
model.config.meanflow = meanflow_config
|
||||
model.core.meanflow_config = meanflow_config
|
||||
model.core.mode = "meanflow"
|
||||
|
||||
old_dit = model.core.velocity_field_predictor
|
||||
if getattr(old_dit, "duration_embedder", None) is not None:
|
||||
return
|
||||
|
||||
new_dit = DiT(
|
||||
in_dim=model.core.fm_hidden_size,
|
||||
out_dim=model.core.latent_dim,
|
||||
transformer_config=model.core.config.DiT,
|
||||
mode="meanflow",
|
||||
)
|
||||
missing_keys, unexpected_keys = new_dit.load_state_dict(
|
||||
old_dit.state_dict(),
|
||||
strict=False,
|
||||
)
|
||||
missing_keys = [
|
||||
key for key in missing_keys if not key.startswith("duration_embedder.")
|
||||
]
|
||||
if missing_keys or unexpected_keys:
|
||||
raise RuntimeError(
|
||||
"Failed to initialize MeanFlow DiT from the pretrained flow-matching "
|
||||
f"DiT: missing={missing_keys[:5]} unexpected={unexpected_keys[:5]}"
|
||||
)
|
||||
duration_output = new_dit.duration_embedder.mlp[-1]
|
||||
nn.init.zeros_(duration_output.weight)
|
||||
nn.init.zeros_(duration_output.bias)
|
||||
model.core.velocity_field_predictor = new_dit
|
||||
|
||||
|
||||
class MeanFlowDotsTtsModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
student: dots_tts_model.DotsTtsModel,
|
||||
settings: MeanFlowSettings,
|
||||
):
|
||||
super().__init__()
|
||||
self.student = student
|
||||
self.settings = settings
|
||||
self._teacher_holder: dict[str, dots_tts_model.DotsTtsModel] = {}
|
||||
|
||||
@property
|
||||
def config(self):
|
||||
return self.student.config
|
||||
|
||||
@property
|
||||
def tokenizer(self):
|
||||
return self.student.tokenizer
|
||||
|
||||
@property
|
||||
def teacher(self) -> dots_tts_model.DotsTtsModel:
|
||||
teacher = self._teacher_holder.get("model")
|
||||
if teacher is None:
|
||||
raise RuntimeError("MeanFlow teacher model has not been initialized.")
|
||||
return teacher
|
||||
|
||||
def set_teacher(self, teacher: dots_tts_model.DotsTtsModel) -> None:
|
||||
for param in teacher.parameters():
|
||||
param.requires_grad_(False)
|
||||
teacher.eval()
|
||||
self._teacher_holder["model"] = teacher
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
super().to(*args, **kwargs)
|
||||
teacher = self._teacher_holder.get("model")
|
||||
if teacher is not None:
|
||||
self._teacher_holder["model"] = teacher.to(*args, **kwargs)
|
||||
self._teacher_holder["model"].eval()
|
||||
return self
|
||||
|
||||
def cuda(self, device=None):
|
||||
super().cuda(device)
|
||||
teacher = self._teacher_holder.get("model")
|
||||
if teacher is not None:
|
||||
self._teacher_holder["model"] = teacher.cuda(device).eval()
|
||||
return self
|
||||
|
||||
def train(self, mode: bool = True):
|
||||
super().train(mode)
|
||||
teacher = self._teacher_holder.get("model")
|
||||
if teacher is not None:
|
||||
teacher.eval()
|
||||
return self
|
||||
|
||||
def prepare_training_batch(self, data: dict[str, Any]) -> dict[str, Any]:
|
||||
return self.student.prepare_training_batch(data)
|
||||
|
||||
def save_pretrained(self, save_directory: str | Path) -> Path:
|
||||
return self.student.save_pretrained(save_directory)
|
||||
|
||||
def load_pretrained_weights(
|
||||
self, pretrained_model_name_or_path: str | Path
|
||||
) -> None:
|
||||
self.student.load_pretrained_weights(pretrained_model_name_or_path)
|
||||
|
||||
def set_cfg_droprate(
|
||||
self,
|
||||
cfg_droprate: float | None = None,
|
||||
xvec_drop_rate: float | None = None,
|
||||
) -> None:
|
||||
self.student.set_cfg_droprate(
|
||||
cfg_droprate=cfg_droprate,
|
||||
xvec_drop_rate=xvec_drop_rate,
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def compute_teacher_meanflow_target(
|
||||
self,
|
||||
*,
|
||||
xt: torch.Tensor,
|
||||
t: torch.Tensor,
|
||||
delta_t: torch.Tensor,
|
||||
prefix_data: dict[str, Any],
|
||||
g_cond: torch.Tensor | None,
|
||||
cfg_distill: bool,
|
||||
uncond_prefix_data: dict[str, Any] | None,
|
||||
uncond_g_cond: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
teacher_core = self.teacher.core
|
||||
teacher_dit = teacher_core.velocity_field_predictor
|
||||
io_helper = teacher_core.io_helper
|
||||
noisy_proj = teacher_core.coordinate_proj
|
||||
n_steps = int(self.settings.teacher_steps)
|
||||
solver = self.settings.teacher_solver
|
||||
cfg_scale = float(self.settings.distill_cfg_scale)
|
||||
|
||||
if solver not in _ALLOWED_TEACHER_SOLVERS:
|
||||
raise ValueError(f"Unsupported teacher solver: {solver!r}.")
|
||||
|
||||
device = xt.device
|
||||
batch_size = xt.size(0)
|
||||
latent_lens = prefix_data["latent_lens"]
|
||||
latent_patch_size = int(prefix_data["latent_patch_size"])
|
||||
anchor_mask = delta_t.float() == 0
|
||||
|
||||
autocast_device = "cuda" if device.type == "cuda" else "cpu"
|
||||
with torch.autocast(device_type=autocast_device, enabled=False):
|
||||
z = xt.float()
|
||||
cur_t = t.float()
|
||||
safe_dt = delta_t.float().clamp(min=1e-6)
|
||||
step_dt = safe_dt / n_steps
|
||||
|
||||
def evaluate(z_in: torch.Tensor, t_val: torch.Tensor) -> torch.Tensor:
|
||||
fm_seq = io_helper.replace_noise_latents_in_fm_seq(
|
||||
prefix_data,
|
||||
z_in.to(xt.dtype),
|
||||
noisy_proj,
|
||||
).float()
|
||||
vt = teacher_dit(
|
||||
x=fm_seq,
|
||||
timesteps=t_val,
|
||||
pos_ids=prefix_data["fm_pos_ids"],
|
||||
mask=prefix_data["fm_seq_mask"],
|
||||
attn_mask=prefix_data["fm_attn_mask"],
|
||||
g_cond=None if g_cond is None else g_cond.float(),
|
||||
)
|
||||
pred = io_helper.get_dit_outputs(
|
||||
pred_v=vt,
|
||||
fm_prefix_lengths=prefix_data["fm_prefix_lengths"],
|
||||
fm_gen_lengths=prefix_data["fm_gen_lengths"],
|
||||
fm_gen_patch_size=prefix_data["fm_gen_patch_size"],
|
||||
latent_patch_size=prefix_data["latent_patch_size"],
|
||||
)
|
||||
|
||||
if cfg_distill:
|
||||
if uncond_prefix_data is None:
|
||||
raise RuntimeError(
|
||||
"CFG distillation requires an uncond prefix."
|
||||
)
|
||||
fm_seq_u = io_helper.replace_noise_latents_in_fm_seq(
|
||||
uncond_prefix_data,
|
||||
z_in.to(xt.dtype),
|
||||
noisy_proj,
|
||||
).float()
|
||||
vt_u = teacher_dit(
|
||||
x=fm_seq_u,
|
||||
timesteps=t_val,
|
||||
pos_ids=uncond_prefix_data["fm_pos_ids"],
|
||||
mask=uncond_prefix_data["fm_seq_mask"],
|
||||
attn_mask=uncond_prefix_data["fm_attn_mask"],
|
||||
g_cond=None if uncond_g_cond is None else uncond_g_cond.float(),
|
||||
)
|
||||
pred_u = io_helper.get_dit_outputs(
|
||||
pred_v=vt_u,
|
||||
fm_prefix_lengths=uncond_prefix_data["fm_prefix_lengths"],
|
||||
fm_gen_lengths=uncond_prefix_data["fm_gen_lengths"],
|
||||
fm_gen_patch_size=uncond_prefix_data["fm_gen_patch_size"],
|
||||
latent_patch_size=uncond_prefix_data["latent_patch_size"],
|
||||
)
|
||||
pred = pred + cfg_scale * (pred - pred_u)
|
||||
return rearrange(pred, "n p d -> (n p) d")
|
||||
|
||||
v_init_flat = evaluate(z, cur_t)
|
||||
|
||||
def apply_velocity(
|
||||
z_cur: torch.Tensor,
|
||||
v_flat: torch.Tensor,
|
||||
*,
|
||||
dt_factor: float,
|
||||
) -> torch.Tensor:
|
||||
new_z = z_cur.clone()
|
||||
offset = 0
|
||||
for batch_idx in range(batch_size):
|
||||
length = int(latent_lens[batch_idx].item())
|
||||
if length <= 0:
|
||||
continue
|
||||
if not bool(anchor_mask[batch_idx].item()):
|
||||
new_z[batch_idx, :length, :] = z_cur[
|
||||
batch_idx, :length, :
|
||||
] + v_flat[offset : offset + length, :] * (
|
||||
step_dt[batch_idx] * float(dt_factor)
|
||||
)
|
||||
offset += length
|
||||
return new_z
|
||||
|
||||
if solver == "euler":
|
||||
v_flat = v_init_flat
|
||||
for step in range(n_steps):
|
||||
if step > 0:
|
||||
v_flat = evaluate(z, cur_t)
|
||||
z = apply_velocity(z, v_flat, dt_factor=1.0)
|
||||
cur_t = cur_t + step_dt
|
||||
elif solver == "midpoint":
|
||||
for step in range(n_steps):
|
||||
k1 = v_init_flat if step == 0 else evaluate(z, cur_t)
|
||||
z_mid = apply_velocity(z, k1, dt_factor=0.5)
|
||||
k2 = evaluate(z_mid, cur_t + 0.5 * step_dt)
|
||||
z = apply_velocity(z, k2, dt_factor=1.0)
|
||||
cur_t = cur_t + step_dt
|
||||
else:
|
||||
for step in range(n_steps):
|
||||
k1 = v_init_flat if step == 0 else evaluate(z, cur_t)
|
||||
z1 = apply_velocity(z, k1, dt_factor=0.5)
|
||||
k2 = evaluate(z1, cur_t + 0.5 * step_dt)
|
||||
z2 = apply_velocity(z, k2, dt_factor=0.5)
|
||||
k3 = evaluate(z2, cur_t + 0.5 * step_dt)
|
||||
z3 = apply_velocity(z, k3, dt_factor=1.0)
|
||||
k4 = evaluate(z3, cur_t + step_dt)
|
||||
z = apply_velocity(
|
||||
z,
|
||||
(k1 + 2.0 * k2 + 2.0 * k3 + k4) / 6.0,
|
||||
dt_factor=1.0,
|
||||
)
|
||||
cur_t = cur_t + step_dt
|
||||
|
||||
mean_velocity = (z - xt.float()) / safe_dt.view(-1, 1, 1)
|
||||
target_chunks = []
|
||||
offset = 0
|
||||
for batch_idx in range(batch_size):
|
||||
length = int(latent_lens[batch_idx].item())
|
||||
if length <= 0:
|
||||
continue
|
||||
if bool(anchor_mask[batch_idx].item()):
|
||||
target_b = v_init_flat[offset : offset + length, :]
|
||||
else:
|
||||
target_b = mean_velocity[batch_idx, :length, :]
|
||||
target_chunks.append(
|
||||
rearrange(target_b, "(n p) d -> n p d", p=latent_patch_size)
|
||||
)
|
||||
offset += length
|
||||
if not target_chunks:
|
||||
raise RuntimeError("Teacher rollout produced no MeanFlow target.")
|
||||
return torch.cat(target_chunks, dim=0).to(xt.dtype)
|
||||
|
||||
def forward(self, data: dict[str, Any]):
|
||||
loss_masks = data["loss_masks"]
|
||||
processed = self.student.prepare_training_inputs(data)
|
||||
processed["input_span_mask"] = data["input_span_mask"]
|
||||
processed["output_span_mask"] = data["output_span_mask"]
|
||||
outputs = self.meanflow_forward(processed)
|
||||
return self.student._compute_loss_terms(
|
||||
outputs,
|
||||
labels=processed["labels"],
|
||||
loss_masks=loss_masks,
|
||||
)
|
||||
|
||||
def meanflow_forward(self, data: dict[str, Any]) -> DotsTtsForwardOutput:
|
||||
core = self.student.core
|
||||
input_ids: torch.Tensor = data["input_ids"]
|
||||
input_ids_lengths: torch.Tensor = data["input_ids_lengths"]
|
||||
input_span_mask: torch.Tensor = data["input_span_mask"]
|
||||
output_span_mask: torch.Tensor = data["output_span_mask"]
|
||||
batch_size = input_ids.size(0)
|
||||
device = input_ids.device
|
||||
|
||||
latents: torch.Tensor | None = data.get("latents")
|
||||
latents_sampled: torch.Tensor | None = data.get("latents_sampled")
|
||||
latent_lengths: torch.Tensor | None = data.get("latent_lengths")
|
||||
has_latents = latents is not None or latents_sampled is not None
|
||||
|
||||
if has_latents:
|
||||
if latents_sampled is None:
|
||||
latents_sampled = core.io_helper.sample_from_latent(latents)
|
||||
patch_embeddings = core.patch_encoder(
|
||||
latents_sampled, x_lens=latent_lengths
|
||||
)
|
||||
valid_patch_counts = latent_lengths // core.latent_patch_size
|
||||
latents_sampled = core.io_helper.normalize(latents_sampled)
|
||||
else:
|
||||
latents_sampled = None
|
||||
patch_embeddings = None
|
||||
valid_patch_counts = torch.zeros(
|
||||
batch_size,
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
|
||||
input_span_counts = input_span_mask.sum(dim=1)
|
||||
if input_span_counts.sum() > 0 and patch_embeddings is None:
|
||||
raise RuntimeError(
|
||||
"Found audio span tokens but no latents provided to compute patch embeddings."
|
||||
)
|
||||
|
||||
inputs_embeds = core.llm.get_input_embeddings()(input_ids)
|
||||
if patch_embeddings is not None:
|
||||
inputs_embeds = inputs_embeds.clone()
|
||||
patch_embeddings = patch_embeddings.to(inputs_embeds.dtype)
|
||||
for batch_idx in range(batch_size):
|
||||
span_num = int(input_span_counts[batch_idx].item())
|
||||
if span_num == 0:
|
||||
continue
|
||||
expected = int(valid_patch_counts[batch_idx].item())
|
||||
if expected != span_num:
|
||||
raise RuntimeError(
|
||||
f"Mismatch between span tokens ({span_num}) and latent patches "
|
||||
f"({expected}) for sample {batch_idx}."
|
||||
)
|
||||
indices = input_span_mask[batch_idx].nonzero(as_tuple=False).squeeze(-1)
|
||||
inputs_embeds[batch_idx, indices, :] = patch_embeddings[
|
||||
batch_idx,
|
||||
:span_num,
|
||||
:,
|
||||
]
|
||||
|
||||
_llm_attn_mask, llm_seq_mask, _ = core.causal_helper.create_causal_mask_and_pos(
|
||||
seq_lens=input_ids_lengths,
|
||||
max_len=input_ids.size(1),
|
||||
)
|
||||
llm_outputs = core.llm(
|
||||
inputs_embeds=inputs_embeds,
|
||||
attention_mask=llm_seq_mask.long(),
|
||||
use_cache=False,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
)
|
||||
llm_logits = llm_outputs.logits
|
||||
llm_hidden = llm_outputs.hidden_states[-1]
|
||||
eos = core.eos_proj(llm_hidden.detach())
|
||||
|
||||
total_patches = int(output_span_mask.sum().item())
|
||||
if total_patches > 0 and latents_sampled is None:
|
||||
raise RuntimeError("MeanFlow training requested but latents are missing.")
|
||||
|
||||
if total_patches > 0:
|
||||
pred, target = self.meanflow_fm_segment(
|
||||
data,
|
||||
llm_hidden=llm_hidden,
|
||||
inputs_embeds=inputs_embeds,
|
||||
output_span_mask=output_span_mask,
|
||||
latents_sampled=latents_sampled,
|
||||
latent_lengths=latent_lengths,
|
||||
)
|
||||
else:
|
||||
pred, target = self.dummy_fm_forward(core, llm_hidden, device)
|
||||
|
||||
return DotsTtsForwardOutput(
|
||||
llm_logits=llm_logits,
|
||||
pred=pred,
|
||||
target=target,
|
||||
eos_out=eos,
|
||||
)
|
||||
|
||||
def meanflow_fm_segment(
|
||||
self,
|
||||
data: dict[str, Any],
|
||||
*,
|
||||
llm_hidden: torch.Tensor,
|
||||
inputs_embeds: torch.Tensor,
|
||||
output_span_mask: torch.Tensor,
|
||||
latents_sampled: torch.Tensor,
|
||||
latent_lengths: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
core = self.student.core
|
||||
teacher_core = self.teacher.core
|
||||
settings = self.settings
|
||||
batch_size = latents_sampled.size(0)
|
||||
device = latents_sampled.device
|
||||
latent_dtype = latents_sampled.dtype
|
||||
first_t = torch.randn(batch_size, device=device, dtype=latent_dtype)
|
||||
second_t = torch.randn(batch_size, device=device, dtype=latent_dtype)
|
||||
first_t = torch.sigmoid(
|
||||
first_t * float(settings.time_sampling_std)
|
||||
+ float(settings.time_sampling_mean)
|
||||
)
|
||||
second_t = torch.sigmoid(
|
||||
second_t * float(settings.time_sampling_std)
|
||||
+ float(settings.time_sampling_mean)
|
||||
)
|
||||
t_vec = torch.minimum(first_t, second_t)
|
||||
delta_t = (first_t - second_t).abs()
|
||||
anchor_mask = torch.rand(batch_size, device=device, dtype=latent_dtype) < float(
|
||||
settings.anchor_prob
|
||||
)
|
||||
delta_t = torch.where(anchor_mask, torch.zeros_like(delta_t), delta_t)
|
||||
z0 = torch.randn_like(latents_sampled)
|
||||
xt = core.fm_helper.sample_x_t(
|
||||
z0,
|
||||
latents_sampled,
|
||||
t_vec.view(-1, 1, 1).to(latent_dtype),
|
||||
)
|
||||
|
||||
fused_cfg = settings.cfg_distill_mode == "fused"
|
||||
if fused_cfg:
|
||||
cfg_mask = torch.zeros(batch_size, device=device, dtype=torch.bool)
|
||||
xvec_drop_mask = torch.zeros(batch_size, device=device, dtype=torch.bool)
|
||||
else:
|
||||
cfg_mask = torch.empty(
|
||||
batch_size, device=device, dtype=torch.float32
|
||||
).uniform_(0, 1) < float(core.cfg_droprate)
|
||||
xvec_drop_mask = torch.empty(
|
||||
batch_size, device=device, dtype=torch.float32
|
||||
).uniform_(0, 1) < float(core.xvec_drop_rate)
|
||||
|
||||
xvec_cond = core.xvec_proj(data["xvector"])
|
||||
vocal_mask = data.get("vocal_mask")
|
||||
if vocal_mask is None:
|
||||
vocal_mask = torch.ones(batch_size, device=device, dtype=torch.bool)
|
||||
xvec_cond = util_module.mask_data(xvec_cond, xvec_drop_mask & vocal_mask)
|
||||
|
||||
hiddens_for_fm = torch.where(
|
||||
output_span_mask.unsqueeze(-1),
|
||||
llm_hidden,
|
||||
inputs_embeds,
|
||||
)
|
||||
prefix_data = core.io_helper.prepare_meanflow_inputs_for_dit(
|
||||
hiddens=hiddens_for_fm,
|
||||
latents=latents_sampled,
|
||||
latent_lens=latent_lengths,
|
||||
hidden_proj=core.hidden_proj,
|
||||
latent_proj=core.latent_proj,
|
||||
noisy_proj=core.coordinate_proj,
|
||||
span_mask=output_span_mask,
|
||||
hidden_patch_size=core.hidden_patch_size,
|
||||
latent_patch_size=core.latent_patch_size,
|
||||
cfg_mask=cfg_mask,
|
||||
noise_latents=xt,
|
||||
)
|
||||
|
||||
uncond_prefix_data = None
|
||||
uncond_g_cond = None
|
||||
with torch.no_grad():
|
||||
teacher_xvec_cond = teacher_core.xvec_proj(data["xvector"])
|
||||
teacher_xvec_cond = util_module.mask_data(
|
||||
teacher_xvec_cond,
|
||||
xvec_drop_mask & vocal_mask,
|
||||
)
|
||||
teacher_prefix_data = (
|
||||
teacher_core.io_helper.prepare_meanflow_inputs_for_dit(
|
||||
hiddens=hiddens_for_fm,
|
||||
latents=latents_sampled,
|
||||
latent_lens=latent_lengths,
|
||||
hidden_proj=teacher_core.hidden_proj,
|
||||
latent_proj=teacher_core.latent_proj,
|
||||
noisy_proj=teacher_core.coordinate_proj,
|
||||
span_mask=output_span_mask,
|
||||
hidden_patch_size=teacher_core.hidden_patch_size,
|
||||
latent_patch_size=teacher_core.latent_patch_size,
|
||||
cfg_mask=cfg_mask,
|
||||
noise_latents=xt,
|
||||
)
|
||||
)
|
||||
if fused_cfg:
|
||||
uncond_prefix_data = (
|
||||
teacher_core.io_helper.prepare_meanflow_inputs_for_dit(
|
||||
hiddens=hiddens_for_fm,
|
||||
latents=latents_sampled,
|
||||
latent_lens=latent_lengths,
|
||||
hidden_proj=teacher_core.hidden_proj,
|
||||
latent_proj=teacher_core.latent_proj,
|
||||
noisy_proj=teacher_core.coordinate_proj,
|
||||
span_mask=output_span_mask,
|
||||
hidden_patch_size=teacher_core.hidden_patch_size,
|
||||
latent_patch_size=teacher_core.latent_patch_size,
|
||||
cfg_mask=torch.ones(
|
||||
batch_size, device=device, dtype=torch.bool
|
||||
),
|
||||
noise_latents=xt,
|
||||
)
|
||||
)
|
||||
uncond_g_cond = torch.zeros_like(teacher_xvec_cond)
|
||||
|
||||
teacher_target = self.compute_teacher_meanflow_target(
|
||||
xt=xt,
|
||||
t=t_vec,
|
||||
delta_t=delta_t,
|
||||
prefix_data=teacher_prefix_data,
|
||||
g_cond=teacher_xvec_cond,
|
||||
cfg_distill=fused_cfg,
|
||||
uncond_prefix_data=uncond_prefix_data,
|
||||
uncond_g_cond=uncond_g_cond,
|
||||
)
|
||||
if anchor_mask.any() and settings.anchor_target == "formula":
|
||||
target = self.replace_anchor_targets_with_formula(
|
||||
teacher_target,
|
||||
z0=z0,
|
||||
latents_sampled=latents_sampled,
|
||||
latent_lengths=latent_lengths,
|
||||
anchor_mask=anchor_mask,
|
||||
)
|
||||
else:
|
||||
target = teacher_target
|
||||
|
||||
student_vt = core.velocity_field_predictor(
|
||||
x=prefix_data["fm_seq"],
|
||||
timesteps=t_vec,
|
||||
duration=delta_t,
|
||||
pos_ids=prefix_data["fm_pos_ids"],
|
||||
mask=prefix_data["fm_seq_mask"],
|
||||
attn_mask=prefix_data["fm_attn_mask"],
|
||||
g_cond=xvec_cond,
|
||||
)
|
||||
pred = core.io_helper.get_dit_outputs(
|
||||
pred_v=student_vt,
|
||||
fm_prefix_lengths=prefix_data["fm_prefix_lengths"],
|
||||
fm_gen_lengths=prefix_data["fm_gen_lengths"],
|
||||
fm_gen_patch_size=prefix_data["fm_gen_patch_size"],
|
||||
latent_patch_size=prefix_data["latent_patch_size"],
|
||||
)
|
||||
return pred, target
|
||||
|
||||
def replace_anchor_targets_with_formula(
|
||||
self,
|
||||
teacher_target: torch.Tensor,
|
||||
*,
|
||||
z0: torch.Tensor,
|
||||
latents_sampled: torch.Tensor,
|
||||
latent_lengths: torch.Tensor,
|
||||
anchor_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
core = self.student.core
|
||||
formula_target = core.fm_helper.compute_u_t(z0, latents_sampled)
|
||||
chunks = []
|
||||
offset = 0
|
||||
for batch_idx in range(latents_sampled.size(0)):
|
||||
length = int(latent_lengths[batch_idx].item())
|
||||
if length <= 0:
|
||||
continue
|
||||
patch_count = length // core.latent_patch_size
|
||||
if bool(anchor_mask[batch_idx].item()):
|
||||
chunks.append(
|
||||
rearrange(
|
||||
formula_target[batch_idx, :length, :],
|
||||
"(n p) d -> n p d",
|
||||
p=core.latent_patch_size,
|
||||
)
|
||||
)
|
||||
else:
|
||||
chunks.append(teacher_target[offset : offset + patch_count])
|
||||
offset += patch_count
|
||||
if not chunks:
|
||||
raise RuntimeError("Anchor target replacement produced no target.")
|
||||
return torch.cat(chunks, dim=0)
|
||||
|
||||
def dummy_fm_forward(
|
||||
self,
|
||||
core,
|
||||
llm_hidden: torch.Tensor,
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
dummy_length = core.latent_patch_size
|
||||
dummy_seq_h = llm_hidden.new_zeros((1, dummy_length, core.llm_hidden_size))
|
||||
dummy_seq_h = core.hidden_proj(dummy_seq_h) * 0.0
|
||||
dummy_seq_l = llm_hidden.new_zeros((1, dummy_length, core.latent_dim))
|
||||
dummy_seq_l = core.latent_proj(dummy_seq_l) * 0.0
|
||||
dummy_seq_c = llm_hidden.new_zeros((1, dummy_length, core.latent_dim))
|
||||
dummy_seq_c = core.coordinate_proj(dummy_seq_c) * 0.0
|
||||
dummy_seq = dummy_seq_h + dummy_seq_l + dummy_seq_c
|
||||
dummy_times = torch.zeros((1,), device=device, dtype=torch.float32)
|
||||
dummy_duration = torch.zeros((1,), device=device, dtype=torch.float32)
|
||||
dummy_attn_mask = torch.ones(
|
||||
(1, dummy_length, dummy_length),
|
||||
device=device,
|
||||
dtype=torch.bool,
|
||||
)
|
||||
dummy_out = core.velocity_field_predictor(
|
||||
x=dummy_seq,
|
||||
timesteps=dummy_times,
|
||||
duration=dummy_duration,
|
||||
attn_mask=dummy_attn_mask,
|
||||
)
|
||||
pred = dummy_out[:, -core.latent_patch_size :, :]
|
||||
return pred, pred.detach()
|
||||
|
||||
|
||||
class DotsTtsMeanFlowTrainingRun(DotsTtsTrainingRun):
|
||||
def __init__(
|
||||
self,
|
||||
cfg: app_config.AppConfig,
|
||||
*,
|
||||
meanflow_settings: MeanFlowSettings,
|
||||
debug_enabled: bool = False,
|
||||
):
|
||||
self.cfg = cfg
|
||||
self.meanflow_settings = meanflow_settings
|
||||
self.progress = train_utils.TrainProgress()
|
||||
self.max_train_steps = int(cfg.train.max_train_steps)
|
||||
self.grad_accumulation_steps = int(cfg.train.gradient_accumulation_steps)
|
||||
self.last_validation_step: int | None = None
|
||||
self.consecutive_empty_epochs = 0
|
||||
self.saved_latest_checkpoint = False
|
||||
self._last_log_step = 0
|
||||
self._last_log_time = 0.0
|
||||
self._debug_enabled = bool(debug_enabled)
|
||||
self._debug_batch_count = 0
|
||||
self._debug_audio_sample_rate = int(self.cfg.train_data.train_audio_sample_rate)
|
||||
|
||||
project_config = ProjectConfiguration(
|
||||
project_dir=self.cfg.train.output_dir,
|
||||
total_limit=self.cfg.train.max_checkpoints_to_keep,
|
||||
)
|
||||
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=False)
|
||||
self.accelerator = Accelerator(
|
||||
kwargs_handlers=[ddp_kwargs],
|
||||
gradient_accumulation_steps=self.grad_accumulation_steps,
|
||||
log_with="tensorboard",
|
||||
project_config=project_config,
|
||||
step_scheduler_with_optimizer=False,
|
||||
)
|
||||
|
||||
util_module.seed_everything(self.cfg.train.seed)
|
||||
|
||||
student = dots_tts_model.DotsTtsModel.from_pretrained(
|
||||
self.cfg.train.pretrained_model_path
|
||||
)
|
||||
student.set_cfg_droprate(
|
||||
cfg_droprate=self.cfg.train.cfg_droprate,
|
||||
xvec_drop_rate=self.cfg.train.xvec_drop_rate,
|
||||
)
|
||||
enable_meanflow_student(student)
|
||||
if not bool(meanflow_settings.train_all_parameters):
|
||||
for param in student.parameters():
|
||||
param.requires_grad_(False)
|
||||
for param in student.core.velocity_field_predictor.parameters():
|
||||
param.requires_grad_(True)
|
||||
model = MeanFlowDotsTtsModel(student, meanflow_settings)
|
||||
|
||||
teacher_path = (
|
||||
meanflow_settings.teacher_model_path or self.cfg.train.pretrained_model_path
|
||||
)
|
||||
teacher = dots_tts_model.DotsTtsModel.from_pretrained(teacher_path)
|
||||
model.set_teacher(teacher)
|
||||
|
||||
optimizer = AdamW(
|
||||
(param for param in model.parameters() if param.requires_grad),
|
||||
lr=self.cfg.train.learning_rate,
|
||||
weight_decay=self.cfg.train.weight_decay,
|
||||
)
|
||||
scheduler = get_cosine_schedule_with_warmup(
|
||||
optimizer,
|
||||
num_warmup_steps=self.cfg.train.warmup_steps,
|
||||
num_training_steps=self.max_train_steps,
|
||||
)
|
||||
self.model, self.optimizer, self.scheduler = self.accelerator.prepare(
|
||||
model,
|
||||
optimizer,
|
||||
scheduler,
|
||||
)
|
||||
self.unwrapped_model = self.accelerator.unwrap_model(self.model)
|
||||
self.unwrapped_model.to(self.accelerator.device)
|
||||
|
||||
expected_sample_rate = int(self.unwrapped_model.config.vocoder.sample_rate)
|
||||
expected_audio_samples_per_llm_token = int(
|
||||
self.unwrapped_model.student.hop_size
|
||||
) * int(self.unwrapped_model.config.patch_size)
|
||||
if int(self.cfg.train_data.train_audio_sample_rate) != expected_sample_rate:
|
||||
raise ValueError(
|
||||
f"train_data.train_audio_sample_rate={int(self.cfg.train_data.train_audio_sample_rate)} "
|
||||
f"does not match the pretrained model sample rate {expected_sample_rate}."
|
||||
)
|
||||
if (
|
||||
int(self.cfg.train_data.audio_samples_per_llm_token)
|
||||
!= expected_audio_samples_per_llm_token
|
||||
):
|
||||
raise ValueError(
|
||||
"train_data.audio_samples_per_llm_token="
|
||||
f"{int(self.cfg.train_data.audio_samples_per_llm_token)} "
|
||||
"does not match the pretrained model audio token contract "
|
||||
f"{expected_audio_samples_per_llm_token}."
|
||||
)
|
||||
if self.cfg.val_data is not None:
|
||||
if int(self.cfg.val_data.train_audio_sample_rate) != expected_sample_rate:
|
||||
raise ValueError(
|
||||
f"val_data.train_audio_sample_rate={int(self.cfg.val_data.train_audio_sample_rate)} "
|
||||
f"does not match the pretrained model sample rate {expected_sample_rate}."
|
||||
)
|
||||
if (
|
||||
int(self.cfg.val_data.audio_samples_per_llm_token)
|
||||
!= expected_audio_samples_per_llm_token
|
||||
):
|
||||
raise ValueError(
|
||||
"val_data.audio_samples_per_llm_token="
|
||||
f"{int(self.cfg.val_data.audio_samples_per_llm_token)} "
|
||||
"does not match the pretrained model audio token contract "
|
||||
f"{expected_audio_samples_per_llm_token}."
|
||||
)
|
||||
|
||||
if self.accelerator.is_main_process:
|
||||
total_params = sum(
|
||||
param.numel() for param in self.unwrapped_model.parameters()
|
||||
)
|
||||
trainable_params = sum(
|
||||
param.numel()
|
||||
for param in self.unwrapped_model.parameters()
|
||||
if param.requires_grad
|
||||
)
|
||||
self.accelerator.print(f"Total parameters: {total_params:,}")
|
||||
self.accelerator.print(f"Trainable parameters: {trainable_params:,}")
|
||||
self.accelerator.print(
|
||||
f"MeanFlow teacher path: {Path(teacher_path).expanduser()}"
|
||||
)
|
||||
self.accelerator.print(
|
||||
f"Distributed type: {self.accelerator.distributed_type}"
|
||||
)
|
||||
|
||||
tokenizer = self.unwrapped_model.tokenizer
|
||||
self.tokenizer = tokenizer
|
||||
train_dataset = data_module.build_training_dataset(
|
||||
self.cfg.train_data,
|
||||
tokenizer=tokenizer,
|
||||
seed=int(self.cfg.train.seed),
|
||||
accelerator=self.accelerator,
|
||||
)
|
||||
self.train_loader = data_module.build_training_dataloader(
|
||||
train_dataset,
|
||||
self.cfg.train_data,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
|
||||
self.val_loader = None
|
||||
if self.cfg.train.eval_interval is not None or self.cfg.train.run_eval_on_start:
|
||||
if self.cfg.val_data is None:
|
||||
raise ValueError(
|
||||
"Validation requires val_data when eval_interval or "
|
||||
"run_eval_on_start is enabled."
|
||||
)
|
||||
validation_data_cfg = self.cfg.val_data.model_copy(deep=True)
|
||||
validation_data_cfg.num_tokens_per_epoch = None
|
||||
val_dataset = data_module.build_validation_dataset(
|
||||
validation_data_cfg,
|
||||
tokenizer=tokenizer,
|
||||
seed=int(self.cfg.train.seed),
|
||||
accelerator=self.accelerator,
|
||||
)
|
||||
self.val_loader = data_module.build_validation_dataloader(
|
||||
val_dataset,
|
||||
validation_data_cfg,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
|
||||
self._resume_if_available()
|
||||
self.train_loader.set_epoch(self.progress.epoch)
|
||||
|
||||
def _write_run_config(self) -> None:
|
||||
if not bool(getattr(self.accelerator, "is_main_process", True)):
|
||||
return
|
||||
output_dir = Path(self.cfg.train.output_dir)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
config_path = output_dir / "config.yml"
|
||||
payload = self.cfg.to_dict()
|
||||
payload["meanflow_train"] = self.meanflow_settings.to_dict()
|
||||
with config_path.open("w", encoding="utf-8") as fout:
|
||||
yaml.safe_dump(
|
||||
payload,
|
||||
fout,
|
||||
sort_keys=False,
|
||||
allow_unicode=True,
|
||||
)
|
||||
|
||||
def _save_checkpoint(self, learning_rate: float) -> None:
|
||||
train_checkpoint.save_train_checkpoint(
|
||||
self.accelerator,
|
||||
self.model,
|
||||
self.optimizer,
|
||||
self.progress,
|
||||
self.cfg.train.output_dir,
|
||||
self.cfg.train.max_checkpoints_to_keep,
|
||||
self.train_loader.state_dict(),
|
||||
{
|
||||
"type": "transformers_cosine_with_warmup_meanflow",
|
||||
"global_step": int(self.progress.global_step),
|
||||
"base_lr": float(self.cfg.train.learning_rate),
|
||||
"current_lr": float(learning_rate),
|
||||
"warmup_steps": int(self.cfg.train.warmup_steps),
|
||||
"max_train_steps": int(self.max_train_steps),
|
||||
"meanflow": self.meanflow_settings.to_dict(),
|
||||
"state_dict": self.scheduler.state_dict(),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def parse_args(argv=None):
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Accelerate MeanFlow distillation entrypoint for dots.tts."
|
||||
)
|
||||
parser.add_argument("--config", default=app_config.DEFAULT_CONFIG_PATH)
|
||||
parser.add_argument(
|
||||
"--debug",
|
||||
action="store_true",
|
||||
help="Print training debug information.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--teacher-model-path",
|
||||
default=None,
|
||||
help=(
|
||||
"Frozen flow-matching teacher model path. Defaults to "
|
||||
"train.pretrained_model_path."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--teacher-steps", type=int, default=8)
|
||||
parser.add_argument(
|
||||
"--teacher-solver",
|
||||
choices=_ALLOWED_TEACHER_SOLVERS,
|
||||
default="euler",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cfg-distill-mode",
|
||||
choices=_ALLOWED_CFG_DISTILL_MODES,
|
||||
default="fused",
|
||||
)
|
||||
parser.add_argument("--distill-cfg-scale", type=float, default=1.2)
|
||||
parser.add_argument("--anchor-prob", type=float, default=0.5)
|
||||
parser.add_argument(
|
||||
"--anchor-target",
|
||||
choices=_ALLOWED_ANCHOR_TARGETS,
|
||||
default="formula",
|
||||
)
|
||||
parser.add_argument("--time-sampling-mean", type=float, default=-0.4)
|
||||
parser.add_argument("--time-sampling-std", type=float, default=1.0)
|
||||
parser.add_argument(
|
||||
"--train-all-parameters",
|
||||
action="store_true",
|
||||
help="Train all regular dots.tts parameters instead of only the DiT.",
|
||||
)
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
args = parse_args(argv)
|
||||
settings = MeanFlowSettings(
|
||||
teacher_model_path=args.teacher_model_path,
|
||||
teacher_steps=args.teacher_steps,
|
||||
teacher_solver=args.teacher_solver,
|
||||
cfg_distill_mode=args.cfg_distill_mode,
|
||||
distill_cfg_scale=args.distill_cfg_scale,
|
||||
anchor_prob=args.anchor_prob,
|
||||
anchor_target=args.anchor_target,
|
||||
time_sampling_mean=args.time_sampling_mean,
|
||||
time_sampling_std=args.time_sampling_std,
|
||||
train_all_parameters=args.train_all_parameters,
|
||||
)
|
||||
return DotsTtsMeanFlowTrainingRun(
|
||||
app_config.load_config(args.config),
|
||||
meanflow_settings=settings,
|
||||
debug_enabled=args.debug,
|
||||
).run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user