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