File size: 5,517 Bytes
9c8f79c 5fda8fd 9c8f79c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 | """
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()
|