first commit

This commit is contained in:
2026-08-15 14:43:56 +08:00
commit 63b8c28e1d
340 changed files with 59515 additions and 0 deletions
+166
View File
@@ -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
View File
@@ -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
+956
View File
@@ -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())