Aratako's picture
Upload folder using huggingface_hub
12d8274 verified
Raw
History Blame Contribute Delete
8.19 kB
import json
from dataclasses import asdict, dataclass, fields
from pathlib import Path
from typing import Any, TypeVar
@dataclass
class ModelConfig:
latent_dim: int = 128
latent_patch_size: int = 1
model_dim: int = 2048
num_layers: int = 24
num_heads: int = 16
mlp_ratio: float = 2.875
text_mlp_ratio: float | None = 2.6
speaker_mlp_ratio: float | None = 2.6
dropout: float = 0.0
text_vocab_size: int = 102400
text_tokenizer_repo: str = "sbintuitions/sarashina2.2-0.5b"
text_add_bos: bool = True
text_dim: int = 1280
text_layers: int = 14
text_heads: int = 10
use_caption_condition: bool = False
caption_vocab_size: int | None = None
caption_tokenizer_repo: str | None = None
caption_add_bos: bool | None = None
caption_dim: int | None = None
caption_layers: int | None = None
caption_heads: int | None = None
caption_mlp_ratio: float | None = None
speaker_dim: int = 1280
speaker_layers: int = 14
speaker_heads: int = 10
speaker_patch_size: int = 1
timestep_embed_dim: int = 512
adaln_rank: int = 256
norm_eps: float = 1e-5
use_duration_predictor: bool = False
duration_aux_dim: int = 14
duration_hidden_dim: int = 1024
duration_layers: int = 3
duration_dropout: float = 0.1
duration_attention_heads: int = 8
duration_architecture: str = "token_sum_adarn_zero_no_aux"
duration_token_init_frames: float = 9.0
duration_speaker_fusion: str = "adarn_zero"
@property
def patched_latent_dim(self) -> int:
return self.latent_dim * self.latent_patch_size
@property
def speaker_patched_latent_dim(self) -> int:
return self.patched_latent_dim * self.speaker_patch_size
@property
def use_speaker_condition(self) -> bool:
# Voice-design checkpoints are caption-driven and intentionally omit
# reference-speaker conditioning to avoid the easier shortcut.
return not bool(self.use_caption_condition)
@property
def text_mlp_ratio_resolved(self) -> float:
if self.text_mlp_ratio is None:
return self.mlp_ratio
return float(self.text_mlp_ratio)
@property
def caption_vocab_size_resolved(self) -> int:
if self.caption_vocab_size is None:
return int(self.text_vocab_size)
return int(self.caption_vocab_size)
@property
def caption_tokenizer_repo_resolved(self) -> str:
if self.caption_tokenizer_repo is None:
return self.text_tokenizer_repo
return str(self.caption_tokenizer_repo)
@property
def caption_add_bos_resolved(self) -> bool:
if self.caption_add_bos is None:
return bool(self.text_add_bos)
return bool(self.caption_add_bos)
@property
def caption_dim_resolved(self) -> int:
if self.caption_dim is None:
return int(self.text_dim)
return int(self.caption_dim)
@property
def caption_layers_resolved(self) -> int:
if self.caption_layers is None:
return int(self.text_layers)
return int(self.caption_layers)
@property
def caption_heads_resolved(self) -> int:
if self.caption_heads is None:
return int(self.text_heads)
return int(self.caption_heads)
@property
def caption_mlp_ratio_resolved(self) -> float:
if self.caption_mlp_ratio is None:
return self.text_mlp_ratio_resolved
return float(self.caption_mlp_ratio)
@property
def speaker_mlp_ratio_resolved(self) -> float:
if self.speaker_mlp_ratio is None:
return self.mlp_ratio
return float(self.speaker_mlp_ratio)
@dataclass
class TrainConfig:
manifest_path: str = ""
output_dir: str = "outputs"
batch_size: int = 8
num_workers: int = 2
dataloader_persistent_workers: bool = False
dataloader_prefetch_factor: int = 2
allow_tf32: bool = False
compile_model: bool = False
train_mode: str = "rf"
learning_rate: float = 1e-4
weight_decay: float = 0.01
optimizer: str = "muon"
adam_beta1: float = 0.9
adam_beta2: float = 0.999
adam_eps: float = 1e-8
muon_momentum: float = 0.95
muon_adjust_lr_fn: str = "match_rms_adamw"
lr_scheduler: str = "none"
warmup_steps: int = 0
caption_warmup: bool = False
caption_warmup_steps: int = 0
stable_steps: int = 0
min_lr_scale: float = 0.1
max_steps: int = 200000
log_every: int = 100
save_every: int = 1000
checkpoint_best_n: int = 0
valid_ratio: float = 0.0
valid_every: int = 0
progress: bool = True
progress_all_ranks: bool = False
precision: str = "bf16"
grad_clip_norm: float = 1.0
gradient_accumulation_steps: int = 1
max_text_len: int = 256
max_caption_len: int | None = None
text_condition_dropout: float = 0.1
caption_condition_dropout: float = 0.1
speaker_condition_dropout: float = 0.1
max_latent_steps: int = 750
fixed_target_latent_steps: int | None = 750
fixed_target_full_mask: bool = True
rf_loss_mode: str = "echo"
duration_loss_weight: float = 0.1
duration_speaker_dropout: float = 0.1
duration_huber_delta: float = 0.1
timestep_logit_mean: float = 0.0
timestep_logit_std: float = 1.0
timestep_stratified: bool = True
timestep_min: float = 0.001
timestep_max: float = 0.999
wandb_enabled: bool = False
wandb_project: str = "Irodori-TTS"
wandb_entity: str | None = None
wandb_run_name: str | None = None
wandb_mode: str = "online"
ddp_find_unused_parameters: bool = False
lora_enabled: bool = False
lora_r: int = 16
lora_alpha: int = 32
lora_dropout: float = 0.0
lora_bias: str = "none"
lora_target_modules: str = "diffusion_attn"
lora_modules_to_save: str | None = "auto"
seed: int = 0
@dataclass
class SamplingConfig:
num_steps: int = 40
cfg_scale_text: float = 3.0
cfg_scale_caption: float = 3.0
cfg_scale_speaker: float = 5.0
cfg_guidance_mode: str = "independent"
cfg_scale: float | None = None
cfg_min_t: float = 0.5
cfg_max_t: float = 1.0
truncation_factor: float | None = None
rescale_k: float | None = None
rescale_sigma: float | None = None
context_kv_cache: bool = True
speaker_kv_scale: float | None = None
speaker_kv_min_t: float | None = 0.9
speaker_kv_max_layers: int | None = None
# Deprecated: inference length is derived from --seconds and codec hop_length.
sequence_length: int | None = None
seed: int = 0
def save_json(path: str | Path, payload: dict) -> None:
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
def dump_configs(path: str | Path, model_cfg: ModelConfig, train_cfg: TrainConfig) -> None:
save_json(path, {"model": asdict(model_cfg), "train": asdict(train_cfg)})
T = TypeVar("T")
def load_experiment_yaml(path: str | Path) -> dict[str, Any]:
"""
Load experiment config YAML. Returns {} for an empty document.
"""
try:
import yaml
except ImportError as exc:
raise RuntimeError(
"PyYAML is required for --config support. Install with `pip install pyyaml`."
) from exc
payload = yaml.safe_load(Path(path).read_text(encoding="utf-8"))
if payload is None:
return {}
if not isinstance(payload, dict):
raise ValueError(f"Config root must be a mapping: {path}")
return payload
def merge_dataclass_overrides(base: T, overrides: dict[str, Any] | None, section: str) -> T:
"""
Merge mapping overrides into a dataclass instance with key validation.
"""
if overrides is None:
return base
if not isinstance(overrides, dict):
raise ValueError(f"Config section '{section}' must be a mapping.")
allowed = {f.name for f in fields(base)}
unknown = sorted(set(overrides) - allowed)
if unknown:
raise ValueError(f"Unknown keys in '{section}' config: {unknown}")
merged = asdict(base)
merged.update(overrides)
return type(base)(**merged)