{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":130287,"databundleVersionId":15633993},{"sourceType":"modelInstanceVersion","sourceId":732859,"databundleVersionId":15477029,"modelInstanceId":558491,"modelId":571058},{"sourceType":"modelInstanceVersion","sourceId":740844,"databundleVersionId":15575334,"modelInstanceId":558490,"modelId":571057}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# %% [markdown]\n# # MOTIONâ€‘S: Improved Textâ€‘toâ€‘Sign Motion Generation (Final, Fixed)\n#\n# **How to use:**\n# 1. Kaggle GPU notebook, add the competition data and the two models.\n# 2. Paste this entire code into one cell.\n# 3. Run all.\n# 4. Submit `submission.csv`.\n\n# %% [code]\n# =============================================================================\n# 1. Install CLIP\n# =============================================================================\n!pip install -q git+https://github.com/openai/CLIP.git\n\n# =============================================================================\n# 2. Imports\n# =============================================================================\nimport os, sys, json, math, random, warnings, gc\nfrom pathlib import Path\nfrom typing import List, Tuple, Dict, Optional, Any\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast, GradScaler\nfrom torch.distributions.categorical import Categorical\n\nimport clip\nfrom tqdm import tqdm\n\nwarnings.filterwarnings('ignore')\n\n# =============================================================================\n# 3. Configuration\n# =============================================================================\nPROJECT_ROOT = Path(\"/kaggle/working\")\nKAGGLE_INPUT = Path(\"/kaggle/input\")\nDATA_BASE = KAGGLE_INPUT / \"competitions/motion-s-hierarchical-text-to-motion-generation-for-sign-language\"\nCSV_PATH = DATA_BASE / \"train.csv\"\nTEST_CSV = DATA_BASE / \"test.csv\"\nVAE_PATH = KAGGLE_INPUT / \"models/antonygithinji/motion-s-vae-rvq/pytorch/default/3/rvq_vae_best.pth\"\nLENGTH_EST_PATH = KAGGLE_INPUT / \"models/antonygithinji/motion-s-length-estimator/pytorch/default/1/length_estimator_best.pth\"\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Using device: {device}\")\n\nVAE_CONFIG = {\n    'num_embeddings': 512,\n    'latent_dim': 256,\n    'num_quantizers': 6,\n}\n\nTRANSFORMER_CONFIG = {\n    'latent_dim': 384,\n    'ff_size': 1024,\n    'num_layers': 8,\n    'num_heads': 6,\n    'dropout': 0.1,\n    'cond_drop_prob': 0.2,\n    'max_token_len': 500,\n    'min_token_len': 6,\n    'text_source': 'both',\n    'batch_size': 32,\n    'grad_accum': 4,\n    'lr': 2e-4,\n    'weight_decay': 0.01,\n    'warmup_epochs': 12,\n    'epochs': 120,\n    'full_mask_prob': 0.4,\n    'label_smoothing': 0.05,\n    'residual_start_epoch': 30,\n    'res_lr': 1e-4,\n    'res_weight_decay': 0.01,\n    'res_prob': 1.0,\n    'save_every': 50,\n    'gen_timesteps': 18,\n    'gen_cond_scale': 5.0,\n    'gen_temperature': 1.0,\n    'gen_topk_thres': 0.9,\n}\n\nTOKEN_COLS = [\"base_tokens\", \"residual_1\", \"residual_2\",\n              \"residual_3\", \"residual_4\", \"residual_5\"]\nVOCAB_SIZE = 512\nMIN_FRAMES, MAX_FRAMES = 40, 800\n\n# =============================================================================\n# 4. Helper functions\n# =============================================================================\ndef set_seed(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nset_seed(42)\n\ndef lens_mask(lengths: torch.Tensor, max_len: int) -> torch.Tensor:\n    device = lengths.device\n    return torch.arange(max_len, device=device).expand(len(lengths), max_len) < lengths.unsqueeze(1)\n\n# =============================================================================\n# 5. Model components (batch_first=True, fixed forward)\n# =============================================================================\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 1000):\n        super().__init__()\n        self.dropout = nn.Dropout(dropout)\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        pe = pe.unsqueeze(0)  # (1, max_len, d_model)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return self.dropout(x + self.pe[:, :x.size(1)])\n\nclass CrossAttentionBlock(nn.Module):\n    def __init__(self, dim: int, heads: int, ff_dim: int = 2048, dropout: float = 0.1):\n        super().__init__()\n        self.self_attn = nn.MultiheadAttention(dim, heads, dropout=dropout, batch_first=True)\n        self.self_attn_norm = nn.LayerNorm(dim)\n        self.self_attn_drop = nn.Dropout(dropout)\n\n        self.cross_attn = nn.MultiheadAttention(dim, heads, dropout=dropout, batch_first=True)\n        self.cross_attn_norm = nn.LayerNorm(dim)\n        self.cross_attn_drop = nn.Dropout(dropout)\n\n        self.ff = nn.Sequential(\n            nn.Linear(dim, ff_dim), nn.GELU(),\n            nn.Dropout(dropout), nn.Linear(ff_dim, dim),\n            nn.Dropout(dropout)\n        )\n        self.ff_norm = nn.LayerNorm(dim)\n\n    def forward(self, x: torch.Tensor, context: torch.Tensor,\n                x_pad: Optional[torch.Tensor] = None,\n                ctx_pad: Optional[torch.Tensor] = None) -> torch.Tensor:\n        # x, context: (batch, seq, dim); masks: (batch, seq) boolean, True = pad\n        # Self attention\n        residual = x\n        x = self.self_attn_norm(x)\n        out, _ = self.self_attn(x, x, x, key_padding_mask=x_pad)\n        x = residual + self.self_attn_drop(out)\n\n        # Cross attention\n        residual = x\n        x = self.cross_attn_norm(x)\n        out, _ = self.cross_attn(x, context, context, key_padding_mask=ctx_pad)\n        x = residual + self.cross_attn_drop(out)\n\n        # FFN\n        residual = x\n        x = self.ff_norm(x)\n        x = residual + self.ff(x)\n        return x\n\nclass CLIPTextEncoder(nn.Module):\n    def __init__(self, version: str = \"ViT-B/32\", device: str = \"cuda\",\n                 freeze: bool = True, max_cache: int = 16000):\n        super().__init__()\n        self.model, _ = clip.load(version, device=device)\n        self.model = self.model.float()\n        self.device = device\n        self.embed_dim = self.model.text_projection.shape[1]\n        if freeze:\n            for p in self.model.parameters():\n                p.requires_grad = False\n            self.model.eval()\n        self._cache = {}\n        self._max_cache = max_cache\n\n    def _warm(self, texts: List[str]):\n        missing = [i for i, t in enumerate(texts) if t not in self._cache]\n        if not missing:\n            return\n        miss_texts = [texts[i] for i in missing]\n        tokens = clip.tokenize(miss_texts, truncate=True).to(self.device)\n        with torch.no_grad():\n            x = self.model.token_embedding(tokens)\n            x += self.model.positional_embedding\n            x = x.permute(1, 0, 2)\n            x = self.model.transformer(x).permute(1, 0, 2)\n            x = self.model.ln_final(x)\n            mask = (tokens != 0).float()\n            e = x[torch.arange(len(miss_texts)), tokens.argmax(-1)] @ self.model.text_projection\n            e = e / e.norm(dim=-1, keepdim=True)\n            for j, i in enumerate(missing):\n                self._cache[texts[i]] = (x[j], mask[j], e[j])\n        if len(self._cache) > self._max_cache:\n            protected = set(texts)\n            to_remove = len(self._cache) - self._max_cache\n            removed = 0\n            for k in list(self._cache.keys()):\n                if removed >= to_remove:\n                    break\n                if k not in protected:\n                    del self._cache[k]\n                    removed += 1\n\n    def encode_text(self, texts: List[str]) -> torch.Tensor:\n        self._warm(texts)\n        return torch.stack([self._cache[t][2] for t in texts])\n\n    def encode_text_tokens(self, texts: List[str]) -> Tuple[torch.Tensor, torch.Tensor]:\n        self._warm(texts)\n        embs = torch.stack([self._cache[t][0] for t in texts])\n        masks = torch.stack([self._cache[t][1] for t in texts])\n        return embs, masks\n\n    def forward(self, texts: List[str], tokens: bool = False):\n        if tokens:\n            return self.encode_text_tokens(texts)\n        return self.encode_text(texts)\n\nclass TextContextualizer(nn.Module):\n    def __init__(self, clip_dim: int = 512, latent_dim: int = 384,\n                 heads: int = 4, layers: int = 2, dropout: float = 0.1):\n        super().__init__()\n        self.proj = nn.Linear(clip_dim, latent_dim)\n        encoder_layer = nn.TransformerEncoderLayer(\n            latent_dim, heads, latent_dim * 4, dropout, 'gelu', batch_first=True\n        )\n        self.encoder = nn.TransformerEncoder(encoder_layer, layers)\n\n    def forward(self, token_embs: torch.Tensor, pad_mask: torch.Tensor) -> torch.Tensor:\n        x = self.proj(token_embs)\n        return self.encoder(x, src_key_padding_mask=pad_mask)\n\nclass MaskTransformer(nn.Module):\n    def __init__(self, num_tokens: int, code_dim: int,\n                 latent_dim: int = 384, ff_size: int = 1024,\n                 num_layers: int = 8, num_heads: int = 6,\n                 dropout: float = 0.1, clip_dim: int = 512,\n                 clip_version: str = \"ViT-B/32\", cond_drop_prob: float = 0.1,\n                 device: str = \"cuda\", max_seq_len: int = 600):\n        super().__init__()\n        self.num_tokens = num_tokens\n        self.latent_dim = latent_dim\n        self.cond_drop_prob = cond_drop_prob\n        self.device = device\n\n        self.mask_id = num_tokens\n        self.pad_id = num_tokens + 1\n\n        self.token_embedding = nn.Embedding(num_tokens + 2, code_dim)\n        self.input_proj = nn.Linear(code_dim, latent_dim)\n        self.output_proj = nn.Sequential(\n            nn.Linear(latent_dim, latent_dim), nn.GELU(),\n            nn.LayerNorm(latent_dim), nn.Linear(latent_dim, num_tokens)\n        )\n        self.pos_encoding = PositionalEncoding(latent_dim, dropout, max_seq_len)\n\n        self.text_encoder = CLIPTextEncoder(clip_version, device, freeze=True)\n        self.text_context = TextContextualizer(clip_dim, latent_dim, dropout=dropout)\n\n        self.blocks = nn.ModuleList([\n            CrossAttentionBlock(latent_dim, num_heads, ff_size, dropout)\n            for _ in range(num_layers)\n        ])\n        self.norm = nn.LayerNorm(latent_dim)\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, (nn.Linear, nn.Embedding)):\n                nn.init.normal_(m.weight, 0, 0.02)\n                if hasattr(m, 'bias') and m.bias is not None:\n                    nn.init.zeros_(m.bias)\n            elif isinstance(m, nn.LayerNorm):\n                nn.init.zeros_(m.bias)\n                nn.init.ones_(m.weight)\n\n    def drop_conditioning(self, cond: torch.Tensor, force_zero: bool = False) -> torch.Tensor:\n        if force_zero:\n            return torch.zeros_like(cond)\n        if not self.training or self.cond_drop_prob == 0:\n            return cond\n        batch_size = cond.shape[0]\n        mask = torch.bernoulli(torch.full((batch_size,), self.cond_drop_prob, device=cond.device))\n        return cond * (1 - mask.view(batch_size, 1, 1))\n\n    def encode_text(self, texts: List[str]) -> Tuple[torch.Tensor, torch.Tensor]:\n        token_embs, pad_mask = self.text_encoder(texts, tokens=True)\n        cond_pad = (pad_mask == 0)          # True for padding\n        context = self.text_context(token_embs, cond_pad)\n        return context, cond_pad\n\n    def forward(self, motion_ids: torch.Tensor, texts: List[str],\n                motion_lengths: torch.Tensor, full_mask_prob: float = 0.5,\n                label_smoothing: float = 0.1) -> Tuple[torch.Tensor, torch.Tensor, float]:\n        batch_size, seq_len = motion_ids.shape\n        valid_mask = lens_mask(motion_lengths, seq_len)\n        pad_mask = ~valid_mask\n\n        # Keep original for accuracy (unchanged by pad_id replacement)\n        original_motion_ids = motion_ids.clone()\n        motion_ids = torch.where(valid_mask, motion_ids, self.pad_id)\n\n        context, cond_pad = self.encode_text(texts)\n        context = self.drop_conditioning(context)\n\n        t = torch.rand(batch_size, device=motion_ids.device)\n        mask_ratio = torch.cos(t * math.pi / 2)\n        num_mask = (seq_len * mask_ratio).round().clamp(min=1)\n        full_mask = torch.bernoulli(torch.full((batch_size,), full_mask_prob, device=motion_ids.device)).bool()\n        num_mask = torch.where(full_mask, motion_lengths.float(), num_mask)\n\n        rand_perm = torch.rand(batch_size, seq_len, device=motion_ids.device).argsort(-1)\n        mask = rand_perm < num_mask.unsqueeze(-1)\n        mask &= valid_mask\n\n        labels = torch.where(mask, motion_ids, self.mask_id)\n        input_ids = motion_ids.clone()\n        rand_mask = torch.bernoulli(torch.full((batch_size, seq_len), 0.1, device=motion_ids.device)).bool() & mask\n        input_ids[rand_mask] = torch.randint(0, self.num_tokens, (batch_size, seq_len), device=motion_ids.device)[rand_mask]\n        mask80 = torch.bernoulli(torch.full((batch_size, seq_len), 0.8, device=motion_ids.device)).bool() & mask & ~rand_mask\n        input_ids[mask80] = self.mask_id\n\n        x = self.token_embedding(input_ids)\n        x = self.input_proj(x)\n        x = self.pos_encoding(x)\n        for blk in self.blocks:\n            x = blk(x, context, pad_mask, cond_pad)\n        x = self.norm(x)\n        logits = self.output_proj(x)          # (batch, seq_len, num_tokens)\n\n        # Loss expects (batch, num_tokens, seq_len)\n        loss = F.cross_entropy(\n            logits.permute(0, 2, 1).reshape(-1, self.num_tokens),\n            labels.reshape(-1),\n            ignore_index=self.mask_id,\n            label_smoothing=label_smoothing\n        )\n        pred = logits.argmax(-1)              # (batch, seq_len)\n\n        # Accuracy only on masked & valid positions\n        valid_masked = mask & valid_mask\n        if valid_masked.any():\n            acc = (pred[valid_masked] == original_motion_ids[valid_masked]).sum().float() / valid_masked.sum().clamp(min=1)\n        else:\n            acc = torch.tensor(0.0, device=motion_ids.device)\n\n        return loss, pred, acc.item()\n\n    def forward_with_cfg(self, token_ids: torch.Tensor, context: torch.Tensor,\n                         cond_pad: torch.Tensor, motion_pad: torch.Tensor,\n                         cond_scale: float = 3.0) -> torch.Tensor:\n        x = self.token_embedding(token_ids)\n        x = self.input_proj(x)\n        x = self.pos_encoding(x)\n        for blk in self.blocks:\n            x = blk(x, context, motion_pad, cond_pad)\n        x = self.norm(x)\n        logits_cond = self.output_proj(x)\n\n        # unconditional\n        x_uncond = self.token_embedding(token_ids)\n        x_uncond = self.input_proj(x_uncond)\n        x_uncond = self.pos_encoding(x_uncond)\n        context_uncond = self.drop_conditioning(context, force_zero=True)\n        for blk in self.blocks:\n            x_uncond = blk(x_uncond, context_uncond, motion_pad, cond_pad)\n        x_uncond = self.norm(x_uncond)\n        logits_uncond = self.output_proj(x_uncond)\n\n        return logits_uncond + cond_scale * (logits_cond - logits_uncond)\n\n    @torch.no_grad()\n    def generate(self, texts: List[str], motion_lengths: torch.Tensor,\n                 timesteps: int = 10, cond_scale: float = 4.0,\n                 temperature: float = 1.0, topk_thres: float = 0.9) -> torch.Tensor:\n        self.eval()\n        batch_size = len(texts)\n        max_len = motion_lengths.max().item()\n        valid_mask = lens_mask(motion_lengths, max_len)\n        pad_mask = ~valid_mask\n\n        context, cond_pad = self.encode_text(texts)\n        token_ids = torch.where(pad_mask, self.pad_id, self.mask_id)\n        confidence = torch.where(pad_mask, 1e5, 0.0)\n\n        for step in range(timesteps):\n            t = torch.tensor(step / timesteps, device=token_ids.device)\n            mask_ratio = torch.cos(t * math.pi / 2)\n            num_mask = (mask_ratio * motion_lengths.float()).round().clamp(min=1).long()\n            ranks = confidence.argsort(1).argsort(1)\n            to_mask = ranks < num_mask.unsqueeze(1)\n            token_ids[to_mask] = self.mask_id\n\n            logits = self.forward_with_cfg(token_ids, context, cond_pad, pad_mask, cond_scale)\n            # logits: (batch, seq_len, num_tokens)\n            k = max(1, int((1 - topk_thres) * logits.shape[-1]))\n            topk_vals, topk_idx = logits.topk(k, dim=-1)\n            filtered = torch.full_like(logits, float('-inf'))\n            filtered.scatter_(-1, topk_idx, topk_vals)\n            probs = F.softmax(filtered / temperature, dim=-1)\n            sampled = torch.multinomial(probs.view(-1, self.num_tokens), 1).view(batch_size, max_len)\n            token_ids = torch.where(to_mask & valid_mask, sampled, token_ids)\n\n            probs_all = F.softmax(logits, dim=-1)\n            confidence = probs_all.gather(2, sampled.unsqueeze(-1)).squeeze(-1)\n            confidence[~to_mask] = 1e5\n\n        token_ids[pad_mask] = -1\n        return token_ids\n\n    def params_no_clip(self):\n        return [p for n, p in self.named_parameters() if 'text_encoder' not in n]\n\n\n# -----------------------------------------------------------------------------\n# ResidualTransformer (unchanged but keep for completeness)\n# -----------------------------------------------------------------------------\nclass ResidualTransformer(nn.Module):\n    def __init__(self, num_tokens: int, code_dim: int, num_quantizers: int,\n                 latent_dim: int = 384, ff_size: int = 1024,\n                 num_layers: int = 8, num_heads: int = 6,\n                 dropout: float = 0.1, clip_dim: int = 512,\n                 clip_version: str = \"ViT-B/32\", cond_drop_prob: float = 0.1,\n                 device: str = \"cuda\", max_seq_len: int = 600,\n                 share_weight: bool = True):\n        super().__init__()\n        self.num_tokens = num_tokens\n        self.num_quantizers = num_quantizers\n        self.latent_dim = latent_dim\n        self.cond_drop_prob = cond_drop_prob\n        self.device = device\n        self.pad_id = num_tokens\n        self.share_weight = share_weight\n\n        if share_weight:\n            self.token_embedding = nn.Embedding(num_tokens + 1, code_dim)\n            self.output_proj = nn.Sequential(\n                nn.Linear(latent_dim, latent_dim), nn.GELU(),\n                nn.LayerNorm(latent_dim), nn.Linear(latent_dim, num_tokens)\n            )\n        else:\n            self.token_embedding = nn.ModuleList([nn.Embedding(num_tokens + 1, code_dim) for _ in range(num_quantizers - 1)])\n            self.output_proj = nn.ModuleList([\n                nn.Sequential(nn.Linear(latent_dim, latent_dim), nn.GELU(),\n                              nn.LayerNorm(latent_dim), nn.Linear(latent_dim, num_tokens))\n                for _ in range(num_quantizers - 1)\n            ])\n\n        self.layer_embedding = nn.Embedding(num_quantizers - 1, latent_dim)\n        self.input_proj = nn.Linear(code_dim, latent_dim)\n        self.pos_encoding = PositionalEncoding(latent_dim, dropout, max_seq_len)\n\n        self.text_encoder = CLIPTextEncoder(clip_version, device, freeze=True)\n        self.text_context = TextContextualizer(clip_dim, latent_dim, dropout=dropout)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            latent_dim, num_heads, ff_size, dropout, 'gelu', batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers)\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, (nn.Linear, nn.Embedding)):\n                nn.init.normal_(m.weight, 0, 0.02)\n                if hasattr(m, 'bias') and m.bias is not None:\n                    nn.init.zeros_(m.bias)\n            elif isinstance(m, nn.LayerNorm):\n                nn.init.zeros_(m.bias)\n                nn.init.ones_(m.weight)\n\n    def drop_conditioning(self, cond: torch.Tensor, force_zero: bool = False) -> torch.Tensor:\n        if force_zero:\n            return torch.zeros_like(cond)\n        if not self.training or self.cond_drop_prob == 0:\n            return cond\n        batch_size = cond.shape[0]\n        mask = torch.bernoulli(torch.full((batch_size,), self.cond_drop_prob, device=cond.device))\n        return cond * (1 - mask.view(batch_size, 1, 1))\n\n    def encode_text(self, texts: List[str]) -> Tuple[torch.Tensor, torch.Tensor]:\n        token_embs, pad_mask = self.text_encoder(texts, tokens=True)\n        cond_pad = (pad_mask == 0)\n        context = self.text_context(token_embs, cond_pad)\n        return context, cond_pad\n\n    def embed_previous_layers(self, prev_tokens: List[torch.Tensor], vq_model) -> torch.Tensor:\n        batch_size, seq_len = prev_tokens[0].shape\n        device = prev_tokens[0].device\n        emb = torch.zeros(batch_size, seq_len, VAE_CONFIG['latent_dim'], device=device)\n        for i, toks in enumerate(prev_tokens):\n            codebook = vq_model.rvq.quantizers[i].embedding\n            flat = toks.reshape(-1).clamp(0, self.num_tokens - 1)\n            layer_emb = codebook[:, flat].t().reshape(batch_size, seq_len, -1)\n            emb += layer_emb\n        return emb\n\n    def forward(self, prev_tokens: List[torch.Tensor], target: torch.Tensor,\n                layer_idx: int, texts: List[str], motion_lengths: torch.Tensor,\n                vq_model) -> Tuple[torch.Tensor, torch.Tensor, float]:\n        batch_size, seq_len = target.shape\n        valid_mask = lens_mask(motion_lengths, seq_len)\n        pad_mask = ~valid_mask\n        target = torch.where(valid_mask, target, self.pad_id)\n\n        prev_emb = self.embed_previous_layers(prev_tokens, vq_model)\n        context, cond_pad = self.encode_text(texts)\n        context = self.drop_conditioning(context)\n\n        x = self.input_proj(prev_emb)\n        layer_tensor = torch.tensor(layer_idx - 1, device=prev_emb.device)\n        x = x + self.layer_embedding(layer_tensor).unsqueeze(0).unsqueeze(1)\n        x = self.pos_encoding(x)\n\n        seq = torch.cat([context, x], dim=1)\n        pad = torch.cat([cond_pad, pad_mask], dim=1)\n        out = self.transformer(seq, src_key_padding_mask=pad)\n        out = out[:, context.size(1):]\n\n        if self.share_weight:\n            logits = self.output_proj(out)\n        else:\n            logits = self.output_proj[layer_idx - 1](out)\n\n        logits = logits.permute(0, 2, 1)\n        loss = F.cross_entropy(logits.reshape(-1, self.num_tokens),\n                               target.reshape(-1),\n                               ignore_index=self.pad_id)\n        pred = logits.argmax(-1)\n        acc = ((pred == target) & valid_mask).sum().float() / valid_mask.sum().clamp(min=1)\n        return loss, pred, acc.item()\n\n    @torch.no_grad()\n    def generate_layer(self, prev_tokens: List[torch.Tensor], layer_idx: int,\n                       texts: List[str], motion_lengths: torch.Tensor,\n                       vq_model, cond_scale: float = 5.0,\n                       temperature: float = 1.0, topk_thres: float = 0.9) -> torch.Tensor:\n        self.eval()\n        batch_size, seq_len = prev_tokens[0].shape[0], prev_tokens[0].shape[1]\n        valid_mask = lens_mask(motion_lengths, seq_len)\n        pad_mask = ~valid_mask\n\n        prev_emb = self.embed_previous_layers(prev_tokens, vq_model)\n        context, cond_pad = self.encode_text(texts)\n\n        def get_logits(cond, force_zero):\n            x = self.input_proj(prev_emb)\n            layer_tensor = torch.tensor(layer_idx - 1, device=prev_emb.device)\n            x = x + self.layer_embedding(layer_tensor).unsqueeze(0).unsqueeze(1)\n            x = self.pos_encoding(x)\n            if force_zero:\n                ctx = self.drop_conditioning(cond, force_zero=True)\n            else:\n                ctx = cond\n            seq = torch.cat([ctx, x], dim=1)\n            pad = torch.cat([cond_pad, pad_mask], dim=1)\n            out = self.transformer(seq, src_key_padding_mask=pad)\n            out = out[:, ctx.size(1):]\n            if self.share_weight:\n                return self.output_proj(out)\n            else:\n                return self.output_proj[layer_idx - 1](out)\n\n        logits = get_logits(context, False) + cond_scale * (get_logits(context, False) - get_logits(context, True))\n\n        k = max(1, int((1 - topk_thres) * logits.shape[-1]))\n        topk_vals, topk_idx = logits.topk(k, dim=-1)\n        filtered = torch.full_like(logits, float('-inf'))\n        filtered.scatter_(-1, topk_idx, topk_vals)\n        probs = F.softmax(filtered / temperature, dim=-1)\n        sampled = torch.multinomial(probs.view(-1, self.num_tokens), 1).view(batch_size, seq_len)\n        sampled = torch.where(valid_mask, sampled, self.pad_id)\n        return sampled\n\n    def params_no_clip(self):\n        return [p for n, p in self.named_parameters() if 'text_encoder' not in n]\n\n\n# =============================================================================\n# 6. VQâ€‘VAE and Length Estimator (unchanged)\n# =============================================================================\nclass VectorQuantizer(nn.Module):\n    def __init__(self, num_embeddings: int = 512, embedding_dim: int = 256):\n        super().__init__()\n        self.num_embeddings = num_embeddings\n        self.embedding_dim = embedding_dim\n        self.embedding = nn.Parameter(torch.randn(embedding_dim, num_embeddings))\n\nclass ResidualVectorQuantizer(nn.Module):\n    def __init__(self, num_quantizers: int = 6, num_embeddings: int = 512, embedding_dim: int = 256):\n        super().__init__()\n        self.num_quantizers = num_quantizers\n        self.quantizers = nn.ModuleList([VectorQuantizer(num_embeddings, embedding_dim) for _ in range(num_quantizers)])\n\nclass RVQVAE(nn.Module):\n    def __init__(self, encoder, decoder, rvq, downsampling_ratio: int = 4):\n        super().__init__()\n        self.encoder = encoder\n        self.decoder = decoder\n        self.rvq = rvq\n        self.downsampling_ratio = downsampling_ratio\n\ndef load_vae(path: str, device: str = 'cuda') -> RVQVAE:\n    ckpt = torch.load(path, map_location=device, weights_only=False)\n    state = ckpt.get('model_state_dict', ckpt.get('state_dict', ckpt))\n    cfg = ckpt.get('model_config', ckpt.get('config', {}))\n    rvq = ResidualVectorQuantizer(\n        num_quantizers=cfg.get('num_quantizers', 6),\n        num_embeddings=cfg.get('num_embeddings', 512),\n        embedding_dim=cfg.get('latent_dim', 256),\n    )\n    rvq_state = {k.replace('rvq.', ''): v for k, v in state.items() if k.startswith('rvq.')}\n    rvq.load_state_dict(rvq_state, strict=False)\n    vae = RVQVAE(None, None, rvq, downsampling_ratio=cfg.get('downsampling_ratio', 4))\n    vae.to(device).eval()\n    print(f\"VAE loaded: {cfg.get('num_quantizers',6)} quantizers, \"\n          f\"{cfg.get('num_embeddings',512)} embeddings, downsampling_ratio={cfg.get('downsampling_ratio',4)}\")\n    return vae\n\nclass LengthEstimator(nn.Module):\n    def __init__(self, clip_dim: int = 512, num_bins: int = 50,\n                 hidden_dim: int = 512, min_tokens: int = 10,\n                 max_tokens: int = 200, dropout: float = 0.2):\n        super().__init__()\n        self.num_bins = num_bins\n        self.min_tokens = min_tokens\n        self.max_tokens = max_tokens\n        self.bin_size = (max_tokens - min_tokens) / (num_bins - 1)\n        self.net = nn.Sequential(\n            nn.Linear(clip_dim, hidden_dim), nn.LayerNorm(hidden_dim),\n            nn.LeakyReLU(0.2, inplace=True), nn.Dropout(dropout),\n            nn.Linear(hidden_dim, hidden_dim // 2), nn.LayerNorm(hidden_dim // 2),\n            nn.LeakyReLU(0.2, inplace=True), nn.Dropout(dropout),\n            nn.Linear(hidden_dim // 2, hidden_dim // 4), nn.LayerNorm(hidden_dim // 4),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Linear(hidden_dim // 4, num_bins)\n        )\n\n    def forward(self, clip_embedding: torch.Tensor) -> torch.Tensor:\n        return self.net(clip_embedding)\n\n    def bin_to_tokens(self, bins: torch.Tensor) -> torch.Tensor:\n        token_len = self.min_tokens + bins.float() * self.bin_size\n        return token_len.round().long()\n\n    @torch.no_grad()\n    def predict(self, clip_embedding: torch.Tensor, mode: str = 'mean',\n                temperature: float = 1.0) -> torch.Tensor:\n        self.eval()\n        logits = self.forward(clip_embedding)\n        if mode == 'argmax':\n            bins = logits.argmax(-1)\n        elif mode == 'sample':\n            probs = F.softmax(logits / temperature, dim=-1)\n            bins = Categorical(probs).sample()\n        else:\n            probs = F.softmax(logits, dim=-1)\n            bin_vals = torch.arange(self.num_bins, device=logits.device, dtype=torch.float)\n            bins = (probs * bin_vals).sum(-1).round().long()\n        return self.bin_to_tokens(bins)\n\n    def estimate_lengths(self, texts: List[str], clip_encoder: CLIPTextEncoder,\n                         mode: str = 'mean', temperature: float = 1.0) -> List[int]:\n        clip_emb = clip_encoder(texts)\n        return self.predict(clip_emb, mode, temperature).tolist()\n\n\n# =============================================================================\n# 7. Data loading and preparation\n# =============================================================================\nclass TokenDataset(Dataset):\n    def __init__(self, data: Dict, text_source: str = 'both',\n                 max_len: int = 500, min_len: int = 6):\n        self.keys = []\n        self.texts = []\n        self.tokens = []\n        self.lengths = []\n        for key, item in data.items():\n            toks = item['tokens']\n            L = toks.shape[1]\n            if L < min_len or L > max_len:\n                continue\n            self.keys.append(key)\n            sentence = str(item.get('sentence', ''))\n            gloss = str(item.get('gloss', ''))\n            if text_source == 'both':\n                txt = f\"{sentence} {gloss}\"\n            elif text_source == 'sentence':\n                txt = sentence\n            else:\n                txt = gloss\n            self.texts.append(txt)\n            padded = np.zeros((toks.shape[0], max_len), dtype=np.int64)\n            padded[:, :L] = toks[:, :L]\n            self.tokens.append(torch.tensor(padded, dtype=torch.long))\n            self.lengths.append(L)\n        print(f\"TokenDataset: {len(self.keys)} samples (filtered length {min_len}-{max_len}), source='{text_source}'\")\n\n    def __len__(self):\n        return len(self.keys)\n\n    def __getitem__(self, idx):\n        return self.texts[idx], self.tokens[idx], self.lengths[idx]\n\ndef collate_fn(batch):\n    texts = [b[0] for b in batch]\n    tokens = torch.stack([b[1] for b in batch])\n    lengths = torch.tensor([b[2] for b in batch], dtype=torch.long)\n    return texts, tokens, lengths\n\ndef load_and_prepare_data(csv_path: str, token_cols: List[str]) -> Dict:\n    df = pd.read_csv(csv_path)\n\n    def parse_tokens(s):\n        if pd.isna(s):\n            return np.array([], dtype=np.int32)\n        return np.array([int(x) for x in str(s).split() if x.strip()], dtype=np.int32)\n\n    for col in token_cols:\n        df[col] = df[col].apply(parse_tokens)\n\n    token_data = {}\n    for _, row in df.iterrows():\n        sid = str(row['id'])\n        base = row['base_tokens']\n        if len(base) == 0:\n            continue\n        canon_len = len(base)\n        layers = []\n        for col in token_cols:\n            t = row[col]\n            if len(t) == 0:\n                t = np.zeros(canon_len, dtype=np.int32)\n            elif len(t) < canon_len:\n                t = np.pad(t, (0, canon_len - len(t)), constant_values=0)\n            else:\n                t = t[:canon_len]\n            layers.append(t)\n        token_data[sid] = {\n            'tokens': np.stack(layers),\n            'sentence': row['sentence'],\n            'gloss': row['gloss'],\n        }\n    return token_data\n\ndef build_retrieval_index(token_data: Dict) -> Dict[str, List[str]]:\n    index = {}\n    for sid, v in token_data.items():\n        g = str(v.get('gloss', '')).strip().lower()\n        if g and g != 'nan':\n            index.setdefault(g, []).append(sid)\n    print(f\"Retrieval index: {len(index)} unique glosses, {sum(len(v) for v in index.values())} samples\")\n    return index\n\n# =============================================================================\n# 8. Training utilities\n# =============================================================================\ndef make_lr_lambda(warmup_epochs: int, total_epochs: int):\n    def lr_lambda(epoch):\n        if epoch < warmup_epochs:\n            return max(0.1, epoch / max(1, warmup_epochs))\n        p = (epoch - warmup_epochs) / max(1, total_epochs - warmup_epochs)\n        return max(0.01, 0.5 * (1 + math.cos(math.pi * p)))\n    return lr_lambda\n\ndef train_epoch(base_model, loader, base_opt, base_sch, device, epoch,\n                res_model=None, res_opt=None, res_sch=None, vq=None,\n                train_res=False, res_prob=1.0, use_amp=True,\n                base_scaler=None, res_scaler=None, grad_accum=4,\n                full_mask_prob=0.5, label_smoothing=0.1):\n    base_model.train()\n    if train_res and res_model:\n        res_model.train()\n\n    total_loss, total_acc, n_batches = 0.0, 0.0, 0\n    res_loss, res_acc, n_res_batches = 0.0, 0.0, 0\n\n    pbar = tqdm(loader, desc=f\"Epoch {epoch}\", leave=False)\n    for i, (texts, tokens, lengths) in enumerate(pbar):\n        tokens = tokens.to(device)\n        lengths = lengths.to(device)\n        base_tokens = tokens[:, 0, :]\n\n        with autocast(device, enabled=use_amp):\n            loss, _, acc = base_model(base_tokens, texts, lengths,\n                                      full_mask_prob, label_smoothing)\n            loss = loss / grad_accum\n        if base_scaler:\n            base_scaler.scale(loss).backward()\n        else:\n            loss.backward()\n\n        if (i + 1) % grad_accum == 0:\n            if base_scaler:\n                base_scaler.step(base_opt)\n                base_scaler.update()\n            else:\n                base_opt.step()\n            base_opt.zero_grad()\n\n        total_loss += loss.item() * grad_accum\n        total_acc += acc\n        n_batches += 1\n\n        if train_res and res_model and random.random() < res_prob:\n            layer_idx = random.randint(1, 5)\n            target = tokens[:, layer_idx, :]\n            prev_tokens = [tokens[:, j, :] for j in range(layer_idx)]\n            with autocast(device, enabled=use_amp):\n                r_loss, _, r_acc = res_model(prev_tokens, target, layer_idx, texts, lengths, vq)\n                r_loss = r_loss / grad_accum\n            if res_scaler:\n                res_scaler.scale(r_loss).backward()\n            else:\n                r_loss.backward()\n            if (i + 1) % grad_accum == 0:\n                if res_scaler:\n                    res_scaler.step(res_opt)\n                    res_scaler.update()\n                else:\n                    res_opt.step()\n                res_opt.zero_grad()\n            res_loss += r_loss.item() * grad_accum\n            res_acc += r_acc\n            n_res_batches += 1\n\n        pbar.set_postfix({\n            'loss': f\"{total_loss / max(1, n_batches):.4f}\",\n            'acc': f\"{total_acc / max(1, n_batches):.4f}\",\n            'r_loss': f\"{res_loss / max(1, n_res_batches):.4f}\" if train_res else 'N/A'\n        })\n\n    if base_sch:\n        base_sch.step()\n    if res_sch:\n        res_sch.step()\n\n    metrics = {'loss': total_loss / max(1, n_batches), 'acc': total_acc / max(1, n_batches)}\n    if train_res:\n        metrics['res_loss'] = res_loss / max(1, n_res_batches)\n        metrics['res_acc'] = res_acc / max(1, n_res_batches)\n    return metrics\n\n@torch.no_grad()\ndef validate(base_model, loader, device, epoch):\n    base_model.eval()\n    total_loss, total_acc, n_batches = 0.0, 0.0, 0\n    for texts, tokens, lengths in tqdm(loader, desc=\"Validation\", leave=False):\n        tokens = tokens.to(device)\n        lengths = lengths.to(device)\n        base_tokens = tokens[:, 0, :]\n        loss, _, acc = base_model(base_tokens, texts, lengths)\n        total_loss += loss.item()\n        total_acc += acc\n        n_batches += 1\n    return {'loss': total_loss / max(1, n_batches), 'acc': total_acc / max(1, n_batches)}\n\n@torch.no_grad()\ndef validate_residual(res_model, loader, device, vq):\n    res_model.eval()\n    total_loss, total_acc, n_batches = 0.0, 0.0, 0\n    for texts, tokens, lengths in tqdm(loader, desc=\"Residual Validation\", leave=False):\n        tokens = tokens.to(device)\n        lengths = lengths.to(device)\n        layer_idx = random.randint(1, 5)\n        target = tokens[:, layer_idx, :]\n        prev_tokens = [tokens[:, j, :] for j in range(layer_idx)]\n        loss, _, acc = res_model(prev_tokens, target, layer_idx, texts, lengths, vq)\n        total_loss += loss.item()\n        total_acc += acc\n        n_batches += 1\n    return {'loss': total_loss / max(1, n_batches), 'acc': total_acc / max(1, n_batches)}\n\n# =============================================================================\n# 9. Main training loop\n# =============================================================================\nprint(\"\\n\" + \"=\"*60)\nprint(\"LOADING DATA AND MODELS\")\nprint(\"=\"*60)\n\ntrain_df = pd.read_csv(CSV_PATH)\ntest_df = pd.read_csv(TEST_CSV)\n\ntoken_data = load_and_prepare_data(CSV_PATH, TOKEN_COLS)\nretrieval_index = build_retrieval_index(token_data)\n\nall_ids = list(token_data.keys())\nrandom.shuffle(all_ids)\nsplit_idx = int(len(all_ids) * 0.9)\ntrain_ids, val_ids = all_ids[:split_idx], all_ids[split_idx:]\ntrain_data = {k: token_data[k] for k in train_ids}\nval_data = {k: token_data[k] for k in val_ids}\n\ncfg = TRANSFORMER_CONFIG\ntrain_dataset = TokenDataset(train_data, cfg['text_source'],\n                             cfg['max_token_len'], cfg['min_token_len'])\nval_dataset = TokenDataset(val_data, cfg['text_source'],\n                           cfg['max_token_len'], cfg['min_token_len'])\n\ntrain_loader = DataLoader(train_dataset, batch_size=cfg['batch_size'], shuffle=True,\n                          collate_fn=collate_fn, pin_memory=True, drop_last=True)\nval_loader = DataLoader(val_dataset, batch_size=cfg['batch_size'], shuffle=False,\n                        collate_fn=collate_fn, pin_memory=True)\n\nvq_model = load_vae(VAE_PATH, device)\n\nmax_seq_len = cfg['max_token_len'] + 77\ntransformer = MaskTransformer(\n    num_tokens=VOCAB_SIZE, code_dim=VAE_CONFIG['latent_dim'],\n    latent_dim=cfg['latent_dim'], ff_size=cfg['ff_size'],\n    num_layers=cfg['num_layers'], num_heads=cfg['num_heads'],\n    dropout=cfg['dropout'], clip_dim=512, clip_version=\"ViT-B/32\",\n    cond_drop_prob=cfg['cond_drop_prob'], device=device, max_seq_len=max_seq_len\n).to(device)\n\nres_model = ResidualTransformer(\n    num_tokens=VOCAB_SIZE, code_dim=VAE_CONFIG['latent_dim'],\n    num_quantizers=VAE_CONFIG['num_quantizers'],\n    latent_dim=cfg['latent_dim'], ff_size=cfg['ff_size'],\n    num_layers=cfg['num_layers'], num_heads=cfg['num_heads'],\n    dropout=cfg['dropout'], clip_dim=512, clip_version=\"ViT-B/32\",\n    cond_drop_prob=cfg['cond_drop_prob'], device=\"cpu\",\n    max_seq_len=max_seq_len, share_weight=True\n)\n\nbase_opt = torch.optim.AdamW(transformer.params_no_clip(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])\nbase_sch = torch.optim.lr_scheduler.LambdaLR(base_opt, make_lr_lambda(cfg['warmup_epochs'], cfg['epochs']))\nuse_amp = True\nbase_scaler = GradScaler(\"cuda\") if use_amp else None\nres_scaler = GradScaler(\"cuda\") if use_amp else None\nres_opt = None\nres_sch = None\n\nprint(f\"\\nMaskTransformer: {sum(p.numel() for p in transformer.parameters()):,} total, \"\n      f\"{sum(p.numel() for p in transformer.params_no_clip()):,} trainable\")\nprint(f\"ResidualTransformer: {sum(p.numel() for p in res_model.parameters()):,} total, \"\n      f\"{sum(p.numel() for p in res_model.params_no_clip()):,} trainable\")\n\nprint(\"\\n\" + \"=\"*60)\nprint(f\"TRAINING: {cfg['epochs']} epochs | Residual starts at epoch {cfg['residual_start_epoch']}\")\nprint(f\"Batch size: {cfg['batch_size']} x {cfg['grad_accum']} = {cfg['batch_size']*cfg['grad_accum']}\")\nprint(f\"full_mask_prob={cfg['full_mask_prob']}, label_smoothing={cfg['label_smoothing']}\")\nprint(\"=\"*60 + \"\\n\")\n\nbest_val_loss = float('inf')\nbest_res_loss = float('inf')\n\nfor epoch in range(1, cfg['epochs'] + 1):\n    train_res = (epoch >= cfg['residual_start_epoch'])\n\n    if train_res and res_opt is None:\n        print(f\"\\n*** Epoch {epoch}: activating ResidualTransformer ***\")\n        torch.cuda.empty_cache()\n        res_model = res_model.to(device)\n        res_model.text_encoder = transformer.text_encoder\n        res_opt = torch.optim.AdamW(res_model.params_no_clip(), lr=cfg['res_lr'], weight_decay=cfg['res_weight_decay'])\n        res_sch = torch.optim.lr_scheduler.LambdaLR(\n            res_opt, make_lr_lambda(cfg['warmup_epochs'], cfg['epochs'] - cfg['residual_start_epoch']))\n\n    train_metrics = train_epoch(\n        transformer, train_loader, base_opt, base_sch, device, epoch,\n        res_model if train_res else None, res_opt, res_sch, vq_model,\n        train_res, cfg['res_prob'], use_amp, base_scaler, res_scaler,\n        cfg['grad_accum'], cfg['full_mask_prob'], cfg['label_smoothing']\n    )\n\n    val_metrics = validate(transformer, val_loader, device, epoch)\n\n    val_res_metrics = {}\n    if train_res and res_opt is not None:\n        val_res_metrics = validate_residual(res_model, val_loader, device, vq_model)\n\n    log = (f\"{epoch:3d}/{cfg['epochs']} | train loss {train_metrics['loss']:.4f} acc {train_metrics['acc']:.3f} | \"\n           f\"val loss {val_metrics['loss']:.4f} acc {val_metrics['acc']:.3f}\")\n    if train_res:\n        log += f\" | res train loss {train_metrics.get('res_loss',0):.4f} acc {train_metrics.get('res_acc',0):.3f}\"\n        log += f\" | res val loss {val_res_metrics.get('loss',0):.4f} acc {val_res_metrics.get('acc',0):.3f}\"\n    print(log)\n\n    if val_metrics['loss'] < best_val_loss:\n        best_val_loss = val_metrics['loss']\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': transformer.state_dict(),\n            'optimizer_state_dict': base_opt.state_dict()\n        }, str(PROJECT_ROOT / \"best_model.pth\"))\n        print(f\"  -> Best base model saved (val_loss={best_val_loss:.4f})\")\n\n    if train_res and val_res_metrics.get('loss', float('inf')) < best_res_loss:\n        best_res_loss = val_res_metrics['loss']\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': transformer.state_dict(),\n            'residual_model_state_dict': res_model.state_dict()\n        }, str(PROJECT_ROOT / \"best_residual_model.pth\"))\n        print(f\"  -> Best residual model saved (val_res_loss={best_res_loss:.4f})\")\n\ntorch.save({\n    'epoch': cfg['epochs'],\n    'model_state_dict': transformer.state_dict(),\n    'residual_model_state_dict': res_model.state_dict()\n}, str(PROJECT_ROOT / \"final_model.pth\"))\nprint(\"\\nTraining complete!\")\n\n# =============================================================================\n# 10. Inference setup\n# =============================================================================\nlength_ckpt = torch.load(LENGTH_EST_PATH, map_location=device, weights_only=False)\nlen_cfg = length_ckpt.get('config', {})\nlength_estimator = LengthEstimator(\n    clip_dim=len_cfg.get('clip_dim', 512),\n    num_bins=len_cfg.get('num_bins', 50),\n    hidden_dim=len_cfg.get('hidden_dim', 512),\n    min_tokens=len_cfg.get('min_tokens', 10),\n    max_tokens=len_cfg.get('max_tokens', 200),\n    dropout=0.0\n)\nif 'model_state_dict' in length_ckpt:\n    length_estimator.load_state_dict(length_ckpt['model_state_dict'])\nelif 'estimator' in length_ckpt:\n    length_estimator.load_state_dict(length_ckpt['estimator'])\nelse:\n    length_estimator.load_state_dict(length_ckpt)\nlength_estimator.to(device).eval()\n\nclip_encoder = transformer.text_encoder\n\ndef estimate_token_length(text: str) -> int:\n    lengths = length_estimator.estimate_lengths([text], clip_encoder, mode='mean', temperature=1.0)\n    return max(10, min(200, lengths[0]))\n\n# =============================================================================\n# 11. Hybrid generation for test set\n# =============================================================================\ndef retrieval_lookup(gloss: str, token_data: Dict, index: Dict) -> Optional[np.ndarray]:\n    g = str(gloss).strip().lower()\n    if g and g != 'nan' and g in index:\n        sid = random.choice(index[g])\n        return token_data[sid]['tokens']\n    return None\n\ndef generate_with_model(text: str, transformer, res_model, vq_model,\n                        length_estimator, clip_encoder, device, cfg) -> List[np.ndarray]:\n    token_len = estimate_token_length(text)\n    motion_lengths = torch.tensor([token_len], device=device)\n    with torch.no_grad():\n        base = transformer.generate(\n            [text], motion_lengths,\n            timesteps=cfg['gen_timesteps'],\n            cond_scale=cfg['gen_cond_scale'],\n            temperature=cfg['gen_temperature'],\n            topk_thres=cfg['gen_topk_thres']\n        )\n        base = torch.clamp(base, 0, VOCAB_SIZE - 1)\n        layers = [base]\n        for li in range(1, 6):\n            res = res_model.generate_layer(\n                layers, li, [text], motion_lengths, vq_model,\n                cond_scale=cfg['gen_cond_scale'],\n                temperature=cfg['gen_temperature'],\n                topk_thres=cfg['gen_topk_thres']\n            )\n            res = torch.clamp(res, 0, VOCAB_SIZE - 1)\n            layers.append(res)\n    return [l[0].cpu().numpy() for l in layers]\n\ndef validate_token_string(token_str: str, target_len: int = None) -> str:\n    tokens = list(map(int, token_str.split()))\n    tokens = [max(0, min(VOCAB_SIZE - 1, t)) for t in tokens if t >= 0]\n    if target_len is not None:\n        if len(tokens) < target_len:\n            last = tokens[-1] if tokens else 0\n            tokens.extend([last] * (target_len - len(tokens)))\n        elif len(tokens) > target_len:\n            tokens = tokens[:target_len]\n    return ' '.join(map(str, tokens))\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"GENERATING PREDICTIONS\")\nprint(\"=\"*60)\n\nresults = []\nretrieval_count, generative_count = 0, 0\n\nfor _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Generating\"):\n    sample_id = row['id']\n    sentence = str(row.get('sentence', ''))\n    gloss = str(row.get('gloss', ''))\n    text = f\"{sentence} {gloss}\"\n\n    retrieved = retrieval_lookup(gloss, token_data, retrieval_index)\n    if retrieved is not None:\n        retrieval_count += 1\n        tokens_list = retrieved\n        L = tokens_list.shape[1]\n        row_dict = {'id': sample_id}\n        for i, col in enumerate(TOKEN_COLS):\n            tok_str = ' '.join(map(str, tokens_list[i].tolist()))\n            row_dict[col] = validate_token_string(tok_str, target_len=L)\n        results.append(row_dict)\n    else:\n        generative_count += 1\n        try:\n            layers = generate_with_model(\n                text, transformer, res_model, vq_model,\n                length_estimator, clip_encoder, device, cfg\n            )\n            L = layers[0].shape[0]\n            row_dict = {'id': sample_id}\n            for i, col in enumerate(TOKEN_COLS):\n                tok_str = ' '.join(map(str, layers[i]))\n                row_dict[col] = validate_token_string(tok_str, target_len=L)\n            results.append(row_dict)\n        except Exception as e:\n            print(f\"  Generation failed for {sample_id}: {e}. Using random fallback.\")\n            fallback_id = random.choice(list(token_data.keys()))\n            fallback_tokens = token_data[fallback_id]['tokens']\n            L = fallback_tokens.shape[1]\n            row_dict = {'id': sample_id}\n            for i, col in enumerate(TOKEN_COLS):\n                tok_str = ' '.join(map(str, fallback_tokens[i].tolist()))\n                row_dict[col] = validate_token_string(tok_str, target_len=L)\n            results.append(row_dict)\n\nprint(f\"\\n  Retrieval matches: {retrieval_count}\")\nprint(f\"  Generated: {generative_count}\")\nprint(f\"  Total: {len(results)}\")\n\nsubmission = pd.DataFrame(results)\nsubmission = submission[['id'] + TOKEN_COLS]\n\nprint(\"\\n--- Submission Validation ---\")\nfor col in TOKEN_COLS:\n    lengths = submission[col].apply(lambda x: len(x.split()))\n    print(f\"  {col}: length range [{lengths.min()}-{lengths.max()}]\")\n\ndef lengths_consistent(row):\n    lens = [len(str(row[c]).split()) for c in TOKEN_COLS]\n    return len(set(lens)) == 1\n\nconsistent = submission.apply(lengths_consistent, axis=1)\nprint(f\"  Consistent lengths across layers: {consistent.all()}\")\nif not consistent.all():\n    print(\"  Fixing inconsistent lengths...\")\n    for idx in submission[~consistent].index:\n        base_len = len(str(submission.loc[idx, 'base_tokens']).split())\n        for col in TOKEN_COLS[1:]:\n            tokens = list(map(int, str(submission.loc[idx, col]).split()))\n            if len(tokens) < base_len:\n                last = tokens[-1] if tokens else 0\n                tokens.extend([last] * (base_len - len(tokens)))\n            elif len(tokens) > base_len:\n                tokens = tokens[:base_len]\n            submission.loc[idx, col] = ' '.join(map(str, tokens))\n\nsubmission.to_csv('submission.csv', index=False)\nprint(\"\\nâœ“ Submission saved to submission.csv\")\nsubmission.head()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}