ViuAI_TTS_200M / scripts /sync_checkpoint.py
ViuAI's picture
Master Fix: Real F0 pitch, acoustic Mel-80 to 100 filterbank vocoder projection, soundfile multi-format loader, safe weights_only
5fda8fd verified
Raw History Blame Contribute Delete
5.52 kB
"""
ViuAI_TTS_200M Checkpoint Synchronization & Migration Script.
Aligns checkpoints/viuai_tts_200m_init.pt with the updated model architecture:
1. Remaps legacy duration_predictor.blocks.* -> prosody_engine.shared_blocks.*
2. Slices vocab embedding from 512 to 224 (preserving all active token embeddings)
3. Initializes missing voice-cloning (dit.ref_audio_proj, dit.null_cond) and pitch/energy heads
4. Verifies 100% strict=True load compatibility (391 tensors, ~190.3M parameters)
"""
import os
import sys
import shutil
import torch
# Ensure local package path is recognized
script_dir = os.path.dirname(os.path.abspath(__file__))
pkg_dir = os.path.dirname(script_dir)
sys.path.insert(0, pkg_dir)
from models import ViuAITTS200M
def sync_checkpoint(ckpt_path: str = None):
if ckpt_path is None:
ckpt_path = os.path.join(pkg_dir, "checkpoints", "viuai_tts_200m_init.pt")
print("=" * 65)
print(f"[*] Synchronizing checkpoint: {ckpt_path}")
print("=" * 65)
if not os.path.exists(ckpt_path):
print(f"[!] Checkpoint not found at {ckpt_path}. Creating a fresh initialized checkpoint.")
model = ViuAITTS200M(vocab_size=224)
new_state = model.state_dict()
else:
# Create backup
backup_path = ckpt_path + ".bak"
if not os.path.exists(backup_path):
shutil.copyfile(ckpt_path, backup_path)
print(f"[+] Created safety backup at: {backup_path}")
try:
raw_ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=True)
except Exception:
raw_ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
old_state = raw_ckpt["model_state_dict"] if "model_state_dict" in raw_ckpt else raw_ckpt
# Instantiate target model
model = ViuAITTS200M(vocab_size=224)
target_state = model.state_dict()
new_state = {}
migrated_count = 0
remapped_count = 0
newly_init_count = 0
for key, target_tensor in target_state.items():
# 1. Exact key match
if key in old_state:
old_tensor = old_state[key]
if old_tensor.shape == target_tensor.shape:
new_state[key] = old_tensor
migrated_count += 1
elif key == "text_encoder.embedding.weight" and old_tensor.shape[0] > target_tensor.shape[0]:
# Slice vocab embedding from 512 to 224
new_state[key] = old_tensor[:target_tensor.shape[0], :].clone()
print(f" [+] Sliced vocab embedding {old_tensor.shape} -> {new_state[key].shape}")
migrated_count += 1
else:
new_state[key] = target_tensor
newly_init_count += 1
# 2. Legacy key mapping for prosody_engine
elif key.startswith("prosody_engine.shared_blocks."):
legacy_key = key.replace("prosody_engine.shared_blocks.", "duration_predictor.blocks.")
if legacy_key in old_state and old_state[legacy_key].shape == target_tensor.shape:
new_state[key] = old_state[legacy_key]
remapped_count += 1
else:
new_state[key] = target_tensor
newly_init_count += 1
elif key == "prosody_engine.duration_head.weight" and "duration_predictor.proj.weight" in old_state:
old_w = old_state["duration_predictor.proj.weight"]
if old_w.shape == target_tensor.shape:
new_state[key] = old_w
remapped_count += 1
else:
new_state[key] = target_tensor
newly_init_count += 1
elif key == "prosody_engine.duration_head.bias" and "duration_predictor.proj.bias" in old_state:
old_b = old_state["duration_predictor.proj.bias"]
if old_b.shape == target_tensor.shape:
new_state[key] = old_b
remapped_count += 1
else:
new_state[key] = target_tensor
newly_init_count += 1
# 3. Missing tensors in old checkpoint (DiT voice cloning, pitch/energy heads)
else:
new_state[key] = target_tensor
newly_init_count += 1
print(f"\n[+] Migration Summary:")
print(f" - Direct Tensor Migrations : {migrated_count}")
print(f" - Remapped Legacy Keys : {remapped_count}")
print(f" - Cleanly Initialized Keys : {newly_init_count}")
print(f" - Total Target Tensors : {len(new_state)}")
# Strict load verification
model.load_state_dict(new_state, strict=True)
print("\n[SUCCESS] strict=True verification passed with 0 missing or unexpected keys!")
# Save re-synchronized checkpoint
torch.save({
"epoch": 0,
"step": 0,
"model_state_dict": model.state_dict(),
"arch": "ViuAITTS200M",
"vocab_size": 224,
"description": "Synchronized initial checkpoint for ViuAI_TTS_200M (Flow Matching + Voice Cloning)",
}, ckpt_path)
total_params = sum(p.numel() for p in model.parameters())
print(f" [SAVED] Checkpoint saved successfully to: {ckpt_path}")
print(f" - Total Parameters: {total_params:,} (~{total_params / 1e6:.2f}M)")
print("=" * 65)
if __name__ == "__main__":
sync_checkpoint()