Download scripts/sync_checkpoint.py from ViuAI/ViuAI_TTS_200M: direct link, hf CLI and curl.
- Browser
- Download file 5.52 kB
-
https://huggingface.co/ViuAI/ViuAI_TTS_200M/resolve/main/scripts/sync_checkpoint.py
- Command line
-
hf download hf://ViuAI/ViuAI_TTS_200M/scripts/sync_checkpoint.py
-
curl -L -o sync_checkpoint.py https://huggingface.co/ViuAI/ViuAI_TTS_200M/resolve/main/scripts/sync_checkpoint.py
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() | |