""" 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()