{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"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":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\n# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:22:16.310459Z","iopub.execute_input":"2026-04-28T14:22:16.310674Z","iopub.status.idle":"2026-04-28T14:24:55.523368Z","shell.execute_reply.started":"2026-04-28T14:22:16.310653Z","shell.execute_reply":"2026-04-28T14:24:55.522384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q git+https://github.com/openai/CLIP.git","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:24:55.525181Z","iopub.execute_input":"2026-04-28T14:24:55.525625Z","iopub.status.idle":"2026-04-28T14:25:04.006925Z","shell.execute_reply.started":"2026-04-28T14:24:55.525593Z","shell.execute_reply":"2026-04-28T14:25:04.006232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport json\nimport math\nimport random\nimport warnings\nimport argparse\nfrom datetime import datetime\nfrom pathlib import Path\nfrom functools import partial\n\n\nimport numpy as np\nimport pandas as pd\n\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.tensorboard import SummaryWriter\nfrom torch.amp import autocast, GradScaler\nimport clip\nimport wandb\n\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\n\nfrom typing import List, Tuple, Dict, Optional","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:04.008785Z","iopub.execute_input":"2026-04-28T14:25:04.009498Z","iopub.status.idle":"2026-04-28T14:25:31.562656Z","shell.execute_reply.started":"2026-04-28T14:25:04.009467Z","shell.execute_reply":"2026-04-28T14:25:31.562090Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PROJECT_ROOT = Path(\"/kaggle/working\")\nKAGGLE_INPUT = Path(\"/kaggle/input\")\nDATA_BASE = KAGGLE_INPUT / \"/kaggle/input/competitions/motion-s-hierarchical-text-to-motion-generation-for-sign-language\"\nCSV_PATH = DATA_BASE / \"train.csv\"\ntest_df = pd.read_csv('/kaggle/input/competitions/motion-s-hierarchical-text-to-motion-generation-for-sign-language/test.csv')\nVAE_PATH = KAGGLE_INPUT / \"/kaggle/input/models/antonygithinji/motion-s-vae-rvq/pytorch/default/3/rvq_vae_best.pth\"\n\nOUTPUT_ROOT = PROJECT_ROOT\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"   Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:31.563710Z","iopub.execute_input":"2026-04-28T14:25:31.564083Z","iopub.status.idle":"2026-04-28T14:25:31.587676Z","shell.execute_reply.started":"2026-04-28T14:25:31.564047Z","shell.execute_reply":"2026-04-28T14:25:31.586971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\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\n    'batch_size': 32,              \n    'grad_accum': 4,                \n    'lr': 2e-4,\n    'weight_decay': 0.01,          \n    'warmup_epochs': 10,          \n    'epochs': 100,                 \n    'full_mask_prob': 0.5,         \n    'label_smoothing': 0.1,        \n    'residual_start_epoch': 50,    \n    'res_lr': 1e-4,\n    'res_weight_decay': 0.01,\n    'res_prob': 1.0,              \n    'save_every': 50,\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:31.588977Z","iopub.execute_input":"2026-04-28T14:25:31.589292Z","iopub.status.idle":"2026-04-28T14:25:31.595120Z","shell.execute_reply.started":"2026-04-28T14:25:31.589268Z","shell.execute_reply":"2026-04-28T14:25:31.594161Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Text Encoding Components\n\nThese modules handle the conversion of input text (glossified sign language) into numerical representations that our motion generation model can process.\n\n**CLIPTextEncoder**: Wraps OpenAI's CLIP text encoder to extract semantic features from glossified text inputs. The encoder can operate in two modes:\n- **Sentence-level**: Returns normalized embeddings for entire text sequences (default)\n- **Token-level**: Returns contextualized embeddings for each token in the sequence (when `tokens=True`)\n\nThe model parameters are frozen by default to preserve CLIP's pre-trained language understanding.\n\n**TextProjector**: A simple feed-forward network that projects CLIP embeddings into the dimensionality expected by downstream motion generation modules. Uses GELU activation and dropout for regularization.","metadata":{}},{"cell_type":"code","source":"class CLIPTextEncoder(nn.Module):\n    \"\"\"\n    CLIP text encoder (frozen).\n    Matches mogen API: encode_text() for global, encode_text_tokens() for per-token.\n    Keeps a string-level cache for speed.\n    \"\"\"\n    def __init__(self, version=\"ViT-B/32\", device=\"cuda\", freeze=True, max_cache=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):\n        \"\"\"Ensure all texts in the batch are cached.\"\"\"\n        miss = [i for i, t in enumerate(texts) if t not in self._cache]\n        if not miss:\n            return\n        mt = [texts[i] for i in miss]\n        tok = clip.tokenize(mt, truncate=True).to(self.device)\n        with torch.no_grad():\n            x = self.model.token_embedding(tok)\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 = (tok != 0).float()\n            e = x[torch.arange(len(mt)), tok.argmax(-1)] @ self.model.text_projection\n            e = e / e.norm(dim=-1, keepdim=True)\n            for j, i in enumerate(miss):\n                self._cache[texts[i]] = (x[j], mask[j], e[j])\n        # Evict oldest entries but PROTECT the current batch\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):\n        \"\"\"Global embeddings (B, embed_dim) — like mogen.\"\"\"\n        self._warm(texts)\n        return torch.stack([self._cache[t][2] for t in texts])\n\n    def encode_text_tokens(self, texts):\n        \"\"\"Token-level embeddings (B, 77, embed_dim) + mask — like mogen.\"\"\"\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, tokens=False):\n        \"\"\"Unified forward (kept for backward compat with mask_trans/res_trans).\"\"\"\n        if tokens:\n            return self.encode_text_tokens(texts)\n        return self.encode_text(texts)\n\n\nclass TextProjector(nn.Module):\n    def __init__(self, din, dout, p=0.1):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(din, dout), nn.GELU(),\n            nn.Dropout(p), nn.Linear(dout, dout)\n        )\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight)\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        return self.net(x)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:31.596497Z","iopub.execute_input":"2026-04-28T14:25:31.596738Z","iopub.status.idle":"2026-04-28T14:25:31.612887Z","shell.execute_reply.started":"2026-04-28T14:25:31.596716Z","shell.execute_reply":"2026-04-28T14:25:31.612189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"clip_text_encoder = CLIPTextEncoder(\"ViT-B/32\", device=device)\ntext_projector = TextProjector(512, TRANSFORMER_CONFIG['latent_dim']).to(device)\n\n# quick test\nwith torch.no_grad():\n    emb = clip_text_encoder([\"Hello, How are you?\"])     # (1, 512)\n    print(\"CLIP emb shape:\", emb.shape)                   # confirm (1, 512)\n    \n    proj = text_projector(emb)\n    print(\"Projected shape:\", proj.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:31.615364Z","iopub.execute_input":"2026-04-28T14:25:31.615714Z","iopub.status.idle":"2026-04-28T14:25:40.008766Z","shell.execute_reply.started":"2026-04-28T14:25:31.615691Z","shell.execute_reply":"2026-04-28T14:25:40.008074Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Mask Transformer","metadata":{}},{"cell_type":"markdown","source":"## Supporting Components and Utilities\n\n**Scheduling and Sampling Utilities**:\n\n- `cos_sched(t)`: Cosine masking schedule that determines what fraction of tokens to mask at each training step. Gradually decreases from 1.0 to 0.0 following a cosine curve, enabling curriculum learning from easier to harder prediction tasks.\n\n- `uniform(shape, dev)`: Simple wrapper for generating random values in [0,1] with proper device placement.\n\n- `topk(logits, th, d)`: Nucleus (top-p) sampling filter that keeps only the most probable tokens whose cumulative probability exceeds the threshold. Prevents sampling from the long tail of unlikely tokens.\n\n- `lens_mask(lens, ml)`: Creates boolean masks for variable-length sequences. Returns a tensor where `True` indicates valid positions and `False` indicates padding positions.\n\n**PositionalEncoding**: Standard sinusoidal position embeddings that inject sequence order information into the transformer. Uses alternating sine and cosine functions at different frequencies to create unique positional patterns.\n\n**CrossAttentionBlock**: The fundamental building block of the transformer, containing three components:\n1. **Self-attention**: Motion tokens attend to other motion tokens in the sequence\n2. **Cross-attention**: Motion tokens attend to text conditioning from CLIP embeddings\n3. **Feed-forward network**: Two-layer MLP with GELU activation for non-linear processing\n\nEach component uses residual connections and layer normalization following the standard transformer architecture.\n\n**Input/Output Processing**:\n\n- `InputProcess`: Projects token embeddings to the model's latent dimension and transposes for sequence-first format required by PyTorch's MultiheadAttention.\n\n- `OutputProcess`: Decodes latent representations back to token logits through a feed-forward layer, layer norm, and final projection to vocabulary size. Transposes output back to batch-first format.","metadata":{}},{"cell_type":"code","source":"def cos_sched(t):\n    return torch.cos(t * math.pi / 2)\n\ndef uniform(shape, dev):\n    return torch.rand(shape, device=dev)\n\ndef topk(logits, th=0.9, d=-1):\n    k = max(1, int((1 - th) * logits.shape[d]))\n    v, i = logits.topk(k, d)\n    out = torch.full_like(logits, float('-inf'))\n    out.scatter_(d, i, v)\n    return out\n\ndef lens_mask(lens, ml):\n    dev = lens.device\n    return torch.arange(ml, device=dev).expand(len(lens), ml) < lens.unsqueeze(1)\n\nclass PositionalEncoding(nn.Module):\n    \"\"\"Sinusoidal positional encoding (seq-first format like mogen).\"\"\"\n    def __init__(self, d, p=0.1, ml=1000):\n        super().__init__()\n        self.drop = nn.Dropout(p)\n        pe = torch.zeros(ml, d)\n        pos = torch.arange(ml).float().unsqueeze(1)\n        div = torch.exp(torch.arange(0, d, 2).float() * (-math.log(10000.0) / d))\n        pe[:, 0::2] = torch.sin(pos * div)\n        pe[:, 1::2] = torch.cos(pos * div)\n        # CRITICAL FIX: Store as (max_len, 1, d_model) like mogen\n        pe = pe.unsqueeze(1)  # (max_len, 1, d_model)\n        self.register_buffer('pe', pe)\n    \n    def forward(self, x):\n        \"\"\"\n        Args:\n            x: (seq_len, B, d_model) - seq-first format\n        Returns:\n            (seq_len, B, d_model) with positional encoding added\n        \"\"\"\n        # x is (seq_len, B, d_model), pe is (max_len, 1, d_model)\n        # pe[:seq_len] broadcasts correctly to (seq_len, B, d_model)\n        return self.drop(x + self.pe[:x.shape[0]])\n\nclass CrossAttentionBlock(nn.Module):\n    def __init__(self, d, h, ff=2048, p=0.1):\n        super().__init__()\n        self.sa = nn.MultiheadAttention(d, h, dropout=p, batch_first=False)\n        self.sa_n = nn.LayerNorm(d)\n        self.sa_p = nn.Dropout(p)\n\n        self.ca = nn.MultiheadAttention(d, h, dropout=p, batch_first=False)\n        self.ca_n = nn.LayerNorm(d)\n        self.ca_p = nn.Dropout(p)\n\n        self.ff = nn.Sequential(\n            nn.Linear(d, ff), nn.GELU(),\n            nn.Dropout(p), nn.Linear(ff, d),\n            nn.Dropout(p)\n        )\n        self.ff_n = nn.LayerNorm(d)\n\n    def forward(self, x, c, xpad=None, cpad=None):\n        r = x; xn = self.sa_n(x)\n        out, _ = self.sa(xn, xn, xn, key_padding_mask=xpad)\n        x = r + self.sa_p(out)\n\n        r = x; xn = self.ca_n(x)\n        out, _ = self.ca(xn, c, c, key_padding_mask=cpad)\n        x = r + self.ca_p(out)\n\n        r = x; xn = self.ff_n(x)\n        x = r + self.ff(xn)\n        return x\n\nclass InputProcess(nn.Module):\n    def __init__(self, cd, ld):\n        super().__init__()\n        self.e = nn.Linear(cd, ld)\n    def forward(self, x):\n        return self.e(x).permute(1, 0, 2)\n\nclass OutputProcess(nn.Module):\n    def __init__(self, ld, nt):\n        super().__init__()\n        self.d = nn.Linear(ld, ld)\n        self.n = nn.LayerNorm(ld)\n        self.o = nn.Linear(ld, nt)\n    def forward(self, x):\n        x = self.d(x)\n        x = F.gelu(x)\n        x = self.n(x)\n        return self.o(x).permute(1, 2, 0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:40.009807Z","iopub.execute_input":"2026-04-28T14:25:40.010227Z","iopub.status.idle":"2026-04-28T14:25:40.024480Z","shell.execute_reply.started":"2026-04-28T14:25:40.010192Z","shell.execute_reply":"2026-04-28T14:25:40.023706Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Masked Transformer Architecture\n\nThis is the core model that generates motion tokens from glossified text inputs. It follows a masked token prediction approach similar to BERT, but adapted for the sequential nature of sign language motion.\n\n**Architecture Overview**:\n- **Text Encoding**: Uses frozen CLIP embeddings projected to the model's latent dimension\n- **Motion Embedding**: Learnable embeddings for each of the 512 possible token values, plus special tokens for masking (`[MASK]`) and padding (`[PAD]`)\n- **Cross-Attention Blocks**: Transformer layers that attend to both the motion sequence and text conditioning\n- **Output Head**: Projects latent representations back to token logits\n\n**Training Strategy**:\nThe model learns through masked token prediction. During training, a portion of tokens are randomly masked using a cosine schedule, and the model must predict the original tokens. This follows the common strategy:\n- 80% of masked positions → replaced with `[MASK]` token\n- 10% of masked positions → replaced with random tokens  \n- 10% of masked positions → kept unchanged\n\n**Generation Process**:\nThe `generate()` method implements iterative refinement over multiple timesteps. Starting from a fully masked sequence, the model progressively unmasks tokens based on prediction confidence, allowing it to build coherent motion sequences that align with the input text. Classifier-free guidance is used to strengthen text-motion alignment.\n\n**Key Parameters**:\n- `cond_scale`: Controls how strongly generation follows text conditioning (higher = more faithful to text)\n- `temperature`: Controls sampling randomness (lower = more deterministic)\n- `topk_filter_thres`: Nucleus sampling threshold for diversity","metadata":{}},{"cell_type":"code","source":"class TextContextualizer(nn.Module):\n    def __init__(self, clip_dim=512, latent_dim=384, heads=4, layers=2, dropout=0.1):\n        super().__init__()\n        self.proj = nn.Linear(clip_dim, latent_dim)\n        layer = nn.TransformerEncoderLayer(\n            latent_dim, heads, latent_dim * 4, \n            dropout, 'gelu', batch_first=True\n        )\n        self.enc = nn.TransformerEncoder(layer, layers)\n    \n    def forward(self, token_embs, pad_mask):\n        # token_embs: (B, 77, 512)  pad_mask: (B, 77) True=ignore\n        x = self.proj(token_embs)              # (B, 77, latent_dim)\n        x = self.enc(x, src_key_padding_mask=pad_mask)\n        return x                               # (B, 77, latent_dim)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:40.025282Z","iopub.execute_input":"2026-04-28T14:25:40.025682Z","iopub.status.idle":"2026-04-28T14:25:40.039819Z","shell.execute_reply.started":"2026-04-28T14:25:40.025660Z","shell.execute_reply":"2026-04-28T14:25:40.039106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass MaskTransformer(nn.Module):\n    def __init__(self, num_tokens, code_dim, latent_dim=384, ff_size=1024, num_layers=8,\n                 num_heads=6, dropout=0.1, clip_dim=512, clip_version=\"ViT-B/32\",\n                 cond_drop_prob=0.1, device=\"cuda\", max_seq_len=600):\n        super().__init__()\n        self.nt = num_tokens\n        self.cd = code_dim\n        self.ld = latent_dim\n        self.cdp = cond_drop_prob\n        self.dev = device\n\n        self.mid = num_tokens\n        self.pid = num_tokens + 1\n\n        self.te = nn.Embedding(num_tokens + 2, code_dim)\n        self.inp = InputProcess(code_dim, latent_dim)\n        self.out = OutputProcess(latent_dim, num_tokens)\n        self.pos = PositionalEncoding(latent_dim, dropout, max_seq_len)\n        self.tcontx = TextContextualizer(clip_dim, latent_dim, dropout=dropout)\n        self.tenc = CLIPTextEncoder(clip_version, device, freeze=True)\n        # self.tproj = TextProjector(clip_dim, latent_dim, dropout)\n\n        self.blks = nn.ModuleList([CrossAttentionBlock(latent_dim, num_heads, ff_size, dropout) for _ in range(num_layers)])\n        self.norm = nn.LayerNorm(latent_dim)\n\n        self.ns = cos_sched\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                m.weight.data.normal_(0, 0.02)\n                if isinstance(m, nn.Linear) and m.bias is not None:\n                    m.bias.data.zero_()\n            elif isinstance(m, nn.LayerNorm):\n                m.bias.data.zero_()\n                m.weight.data.fill_(1.0)\n\n    def dropc(self, c, force=False):\n        if force: return torch.zeros_like(c)\n        if not self.training or self.cdp == 0: return c\n        bs = c.shape[0]\n        m = torch.bernoulli(torch.full((bs,), self.cdp, device=c.device))\n        return c * (1 - m.view(bs, 1, 1))\n\n    def enct(self, texts):\n        te, tm = self.tenc(texts, tokens=True)   # te: (B, 77, 512), tm: (B, 77)\n        cpad = (tm == 0)                          # True = padding position\n        ce = self.tcontx(te, cpad)               # (B, 77, latent_dim) — contextually enriched\n        return ce, cpad\n\n    def fwd_trans(self, mids, ce, cpad, mpad):\n        B, sl = mids.shape\n        x = self.te(mids)\n        x = self.inp(x)\n        x = self.pos(x)\n        text = ce.permute(1, 0, 2)\n\n        for b in self.blks:\n            x = b(x, text, mpad, cpad)\n\n        x = self.norm(x)\n        return self.out(x)\n\n    def forward(self, motion_ids, texts, m_lens, full_mask_prob=0.5, label_smoothing=0.1):\n        B, sl = motion_ids.shape\n        dev = motion_ids.device\n\n        valid = lens_mask(m_lens, sl)\n        pad = ~valid\n        motion_ids = torch.where(valid, motion_ids, self.pid)\n\n        ce, cpad = self.enct(texts)\n        ce = self.dropc(ce)\n\n        t = uniform((B,), dev)\n        mp = self.ns(t)\n        nm = (sl * mp).round().clamp(min=1)\n\n        full = torch.bernoulli(torch.full((B,), full_mask_prob, device=dev)).bool()\n        nm = torch.where(full, m_lens.float(), nm)\n\n        perm = torch.rand((B, sl), device=dev).argsort(-1)\n        mask = perm < nm.unsqueeze(-1)\n        mask &= valid\n\n        lbl = torch.where(mask, motion_ids, self.mid)\n\n        xids = motion_ids.clone()\n        r10 = torch.bernoulli(torch.full((B, sl), 0.1, device=dev)).bool() & mask\n        xids[r10] = torch.randint(0, self.nt, (B, sl), device=dev)[r10]\n\n        m80 = torch.bernoulli(torch.full((B, sl), 0.8, device=dev)).bool() & mask & ~r10\n        xids[m80] = self.mid\n\n        logits = self.fwd_trans(xids, ce, cpad, pad)\n        logits = logits.permute(0, 2, 1)\n\n        loss = F.cross_entropy(\n            logits.reshape(-1, self.nt),\n            lbl.reshape(-1),\n            ignore_index=self.mid,\n            label_smoothing=label_smoothing\n        )\n\n        pred = logits.argmax(-1)\n        acc = ((pred == motion_ids) & mask).sum().float() / mask.sum().clamp(min=1)\n\n        return loss, pred, acc.item()\n\n    def forward_with_cfg(self, motion_ids, cond_emb, cond_padding_mask, motion_padding_mask, cond_scale=3.0):\n        lc = self.fwd_trans(motion_ids, cond_emb, cond_padding_mask, motion_padding_mask)\n        lu = self.fwd_trans(motion_ids, self.dropc(cond_emb, True), cond_padding_mask, motion_padding_mask)\n        return lu + cond_scale * (lc - lu)\n\n    @torch.no_grad()\n    def generate(self, texts, m_lens, timesteps=10, cond_scale=4.0, temperature=1.0, topk_filter_thres=0.9):\n        self.eval()\n        B = len(texts)\n        dev = next(self.parameters()).device\n        ml = m_lens.max().item()\n\n        valid = lens_mask(m_lens, ml)\n        pad = ~valid\n\n        ce, cpad = self.enct(texts)\n\n        ids = torch.where(pad, self.pid, self.mid)\n        conf = torch.where(pad, 1e5, 0.)\n\n        for s in range(timesteps):\n            t = torch.tensor(s / timesteps, device=dev)\n            p = self.ns(t)\n            nm = (p * m_lens.float()).round().clamp(min=1).long()\n\n            ranks = conf.argsort(1).argsort(1)\n            masked = ranks < nm.unsqueeze(1)\n            ids[masked] = self.mid\n\n            logits = self.forward_with_cfg(ids, ce, cpad, pad, cond_scale)\n            logits = logits.permute(0, 2, 1)\n\n            filt = topk(logits, topk_filter_thres)\n            probs = F.softmax(filt / temperature, -1)\n            samp = torch.multinomial(probs.view(-1, self.nt), 1).view(B, ml)\n\n            ids = torch.where(masked & valid, samp, ids)\n\n            pr = F.softmax(logits, -1)\n            conf = pr.gather(2, samp.unsqueeze(-1)).squeeze(-1)\n            conf[~masked] = 1e5\n\n        ids[pad] = -1\n        return ids\n\n    def params_no_clip(self):\n        return [p for n,p in self.named_parameters() if 'tenc' not in n]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:40.041018Z","iopub.execute_input":"2026-04-28T14:25:40.041332Z","iopub.status.idle":"2026-04-28T14:25:40.064877Z","shell.execute_reply.started":"2026-04-28T14:25:40.041310Z","shell.execute_reply":"2026-04-28T14:25:40.063970Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Residual Transfromers","metadata":{}},{"cell_type":"markdown","source":"## Residual Transformer Architecture\n\nThis model generates the residual quantizer layers (layers 1-5) that refine the coarse motion captured by the base layer. Unlike the base model which starts from masked tokens, this model operates autoregressively on the RVQ hierarchy.\n\n**Key Differences from Base Model**:\n\n- **Hierarchical Conditioning**: Takes all previously generated layers as input and predicts the next refinement layer\n- **VAE Codebook Embeddings**: Directly embeds previous layer tokens using the frozen VAE's learned codebooks, preserving semantic motion structure\n- **Layer-Specific Generation**: Processes one residual layer at a time (layer index `li` determines which layer to generate)\n- **Standard Transformer Encoder**: Uses self-attention only (no cross-attention blocks), with text and motion tokens concatenated into a single sequence\n\n**Architecture Flow**:\n1. Embed all previous layers using VAE codebooks and sum them\n2. Add layer-specific embeddings to indicate which residual level to predict\n3. Concatenate text conditioning with motion embeddings\n4. Process through transformer encoder\n5. Decode to token logits for the target layer\n\n**Weight Sharing**:\nThe `share_weight` parameter controls whether all residual layers use shared embeddings and output heads (more parameter-efficient) or have layer-specific parameters (more expressive).\n\n**Generation**:\nThe `generate_layer()` method produces one residual layer at a time using classifier-free guidance. Each layer refines the motion representation, progressively adding finer details to the sign language animation.","metadata":{}},{"cell_type":"code","source":"class ResidualTransformer(nn.Module):\n    def __init__(self, num_tokens, code_dim, num_quantizers, latent_dim=384, ff_size=1024,\n                 num_layers=8, num_heads=6, dropout=0.1, clip_dim=512, clip_version=\"ViT-B/32\",\n                 cond_drop_prob=0.1, device=\"cuda\", max_seq_len=600, share_weight=True):\n        super().__init__()\n        self.nt = num_tokens\n        self.cd = code_dim\n        self.nq = num_quantizers\n        self.ld = latent_dim\n        self.cdp = cond_drop_prob\n        self.dev = device\n        self.share = share_weight\n\n        self.pid = num_tokens\n\n        if share_weight:\n            self.te = nn.Embedding(num_tokens + 1, code_dim)\n            self.out = OutputProcess(latent_dim, num_tokens)\n        else:\n            self.te = nn.ModuleList([nn.Embedding(num_tokens + 1, code_dim) for _ in range(num_quantizers - 1)])\n            self.out = nn.ModuleList([OutputProcess(latent_dim, num_tokens) for _ in range(num_quantizers - 1)])\n\n        self.le = nn.Embedding(num_quantizers - 1, latent_dim)\n        self.inp = InputProcess(code_dim, latent_dim)\n        self.pos = PositionalEncoding(latent_dim, dropout, max_seq_len)\n\n        self.tenc = CLIPTextEncoder(clip_version, device, freeze=True)\n        self.tcontx = TextContextualizer(clip_dim, latent_dim, dropout=dropout)\n\n        enc_l = nn.TransformerEncoderLayer(latent_dim, num_heads, ff_size, dropout, 'gelu')\n        self.trans = nn.TransformerEncoder(enc_l, num_layers)\n\n        self._init_weights()   # ← safe call after creation\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, (nn.Linear, nn.Embedding)):\n                m.weight.data.normal_(0, 0.02)\n                if isinstance(m, nn.Linear) and m.bias is not None:\n                    m.bias.data.zero_()   # safe here (not during apply recursion)\n            elif isinstance(m, nn.LayerNorm):\n                m.bias.data.zero_()\n                m.weight.data.fill_(1.0)\n\n    def dropc(self, c, force=False):\n        if force: return torch.zeros_like(c)\n        if not self.training or self.cdp == 0: return c\n        bs = c.shape[0]\n        m = torch.bernoulli(torch.full((bs,), self.cdp, device=c.device))\n        return c * (1 - m.view(bs, 1, 1))\n\n    # def enct(self, texts):\n    #     te, tm = self.tenc(texts, tokens=True)\n    #     return self.tproj(te), (tm == 0)\n\n\n    def enct(self, texts):\n        te, tm = self.tenc(texts, tokens=True)\n        cpad = (tm == 0)\n        ce = self.tcontx(te, cpad)\n        return ce, cpad\n\n    def embed_prev(self, prev_toks, vq):\n        B, sl = prev_toks[0].shape\n        dev = prev_toks[0].device\n        emb = torch.zeros(B, sl, self.cd, device=dev)\n\n        for i, toks in enumerate(prev_toks):\n            cb = vq.rvq.quantizers[i].embedding\n            flat = toks.reshape(-1).clamp(0, self.nt - 1)\n            le = cb[:, flat].t().reshape(B, sl, self.cd)\n            emb += le\n\n        return emb\n\n    def fwd_trans(self, p_emb, li, ce, cpad, mpad):\n        x = self.inp(p_emb)                              # (sl, B, dim)\n\n        li_idx = torch.tensor(li - 1, device=p_emb.device)\n        x = x + self.le(li_idx).unsqueeze(0).unsqueeze(1) # (sl, B, dim)\n\n        cond = ce.permute(1, 0, 2)                        # (tl, B, dim)\n        seq = torch.cat([cond, x], dim=0)                  # (tl+sl, B, dim)\n        seq = self.pos(seq)\n        pm = torch.cat([cpad, mpad], dim=1)\n\n        out = self.trans(seq, src_key_padding_mask=pm)     # (tl+sl, B, dim)\n        out = out[ce.shape[1]:]                            # (sl, B, dim)\n\n        if self.share:\n            return self.out(out)\n        return self.out[li - 1](out)\n\n    def forward(self, prev_toks, targ, li, texts, mlens, vq):\n        B, sl = targ.shape\n        dev = targ.device\n\n        valid = lens_mask(mlens, sl)\n        pad = ~valid\n        targ = torch.where(valid, targ, self.pid)\n\n        p_emb = self.embed_prev(prev_toks, vq)\n\n        ce, cpad = self.enct(texts)\n        ce = self.dropc(ce)\n\n        logits = self.fwd_trans(p_emb, li, ce, cpad, pad)\n        logits = logits.permute(0, 2, 1)\n\n        loss = F.cross_entropy(\n            logits.reshape(-1, self.nt),\n            targ.reshape(-1),\n            ignore_index=self.pid\n        )\n\n        pred = logits.argmax(-1)\n        acc = ((pred == targ) & valid).sum().float() / valid.sum().clamp(min=1)\n\n        return loss, pred, acc.item()\n\n    def fwd_cfg(self, p_emb, li, ce, cpad, mpad, cs=3.0):\n        lc = self.fwd_trans(p_emb, li, ce, cpad, mpad)\n        lu = self.fwd_trans(p_emb, li, self.dropc(ce, True), cpad, mpad)\n        return lu + cs * (lc - lu)\n\n    @torch.no_grad()\n    def generate_layer(self, prev_toks, li, texts, mlens, vq, cs=5.0, temp=1.0, th=0.9):\n        self.eval()\n        B = len(texts)\n        dev = prev_toks[0].device\n        sl = prev_toks[0].shape[1]\n\n        valid = lens_mask(mlens, sl)\n        pad = ~valid\n\n        p_emb = self.embed_prev(prev_toks, vq)\n\n        ce, cpad = self.enct(texts)\n\n        lc = self.fwd_trans(p_emb, li, ce, cpad, pad)\n        lu = self.fwd_trans(p_emb, li, self.dropc(ce, True), cpad, pad)\n        logits = lu + cs * (lc - lu)\n\n        logits = logits.permute(0, 2, 1)\n\n        k = max(1, int((1 - th) * logits.shape[-1]))\n        tv, ti = logits.topk(k, -1)\n        filt = torch.full_like(logits, float('-inf'))\n        filt.scatter_(-1, ti, tv)\n\n        probs = F.softmax(filt / temp, -1)\n        ids = torch.multinomial(probs.view(-1, self.nt), 1).view(B, sl)\n\n        ids = torch.where(valid, ids, self.pid)\n        return ids\n\n    def params_no_clip(self):\n        return [p for n,p in self.named_parameters() if 'tenc' not in n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:40.065976Z","iopub.execute_input":"2026-04-28T14:25:40.066268Z","iopub.status.idle":"2026-04-28T14:25:40.087854Z","shell.execute_reply.started":"2026-04-28T14:25:40.066236Z","shell.execute_reply":"2026-04-28T14:25:40.087163Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Mask & Residual Transformers","metadata":{}},{"cell_type":"markdown","source":"## Training Data Pipeline and Loop\n\n**Set Seed**: Ensures reproducible results by fixing random seeds across all libraries (Python, NumPy, PyTorch, CUDA).\n\n**TokenDataset**: Loads pre-tokenized motion sequences from the training data with flexible text conditioning options:\n- `src='sentence'`: Use natural English only\n- `src='gloss'`: Use glossified text only  \n- `src='both'`: Concatenate both (default format: \"Sentence: ... Signs: ...\")\n- `src='random'`: Randomly choose between sentence/gloss per sample\n\nHandles variable-length sequences through padding to `max_len` and filters samples by minimum/maximum length constraints.\n\n**Collate Function**: Batches samples together, stacking tokens and lengths into tensors while keeping text as a list of strings for CLIP encoding.\n\n**Training Loop** (`train_epoch`): Implements the two-stage training pipeline:\n\n1. **Base Model Training**: Trains the masked transformer to generate layer 0 (coarse motion) using masked token prediction\n2. **Residual Model Training**: Probabilistically trains residual layers (1-5) to refine the base motion\n\nKey features:\n- **Gradient Accumulation**: Enables larger effective batch sizes on limited GPU memory\n- **Mixed Precision (AMP)**: Optional automatic mixed precision training for faster computation\n- **Curriculum Learning**: `full_mask_prob` controls how often to mask entire sequences (harder task)\n- **Label Smoothing**: Prevents overconfident predictions\n- **Joint Training**: Can train both models simultaneously or independently\n\n**Validation** (`validate`): Evaluates model performance on held-out data. Optionally generates full motion sequences and computes reconstruction metrics (MSE, MAE, RMSE) by decoding through the VAE and comparing to ground truth in motion space.","metadata":{}},{"cell_type":"code","source":"# def init_wandb(config=None, run_name=None):\n#     wandb.init(project='kaggle_benchmarks', name=run_name, config=config or {})\n#     return wandb.run\n\ndef set_seed(s=42):\n    random.seed(s)\n    np.random.seed(s)\n    torch.manual_seed(s)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(s)\n\n\ndef make_lr_lambda(warmup_epochs, total_epochs):\n    \"\"\"Warmup + cosine decay LR schedule (matches mogen).\"\"\"\n    def lr_lambda(epoch):\n        if epoch < warmup_epochs:\n            return max((epoch + 1) / (warmup_epochs + 1), 0.01)\n        progress = (epoch - warmup_epochs) / max(total_epochs - warmup_epochs, 1)\n        return 0.5 * (1 + np.cos(np.pi * progress))\n    return lr_lambda\n\n\nclass TokenDataset(Dataset):\n    COMB = \"Sentence: {sentence} Signs: {gloss}\"\n\n    def __init__(self, data, src='both', ml=80, minl=6, alll=True):\n        self.src = src\n        self.ml = ml\n        self.alll = alll\n\n        self.samples = []\n        for sid, d in data.items():\n            t = d['tokens']\n            if len(t.shape) == 1:\n                t = t[np.newaxis, :]\n            nq, sl = t.shape\n            if minl <= sl <= ml:\n                self.samples.append({\n                    'id': sid,\n                    't': t,\n                    's': d['sentence'],\n                    'g': d['gloss'],\n                    'l': sl,\n                    'nq': nq\n                })\n\n        print(f\"TokenDataset: {len(self.samples)} samples \"\n              f\"(filtered {len(data)} -> len {minl}-{ml}), src='{src}'\")\n\n    def __len__(self):\n        return len(self.samples)\n\n    def _txt(self, item):\n        s, g = item['s'], item['g']\n        if self.src == 'sentence': return s\n        if self.src == 'gloss': return g\n        if self.src == 'both': return self.COMB.format(sentence=s, gloss=g)\n        if self.src == 'random':\n            return s if random.random() < 0.5 else g\n        return s\n\n    def __getitem__(self, i):\n        item = self.samples[i]\n        t = item['t'].copy()\n        txt = self._txt(item)\n        l = item['l']\n\n        if t.shape[1] < self.ml:\n            pad = np.zeros((t.shape[0], self.ml - t.shape[1]), dtype=t.dtype)\n            t = np.concatenate([t, pad], 1)\n        else:\n            t = t[:, :self.ml]\n            l = self.ml\n\n        t = torch.from_numpy(t).long()\n\n        return txt, t if self.alll else t[0], l\n\n\ndef collate(batch, alll=True):\n    txts, ts, ls = zip(*batch)\n    ls = torch.tensor(ls, dtype=torch.long)\n    ts = torch.stack(ts)\n    return list(txts), ts, ls\n\n\ndef train_epoch(\n    base_m, loader, base_opt, base_sch, dev, ep,\n    writer=None, res_m=None, res_opt=None, res_sch=None,\n    vq=None, train_res=False, res_p=0.5, amp=True, base_scaler=None,res_scaler=None,\n    accum=1, res_only=False, fmp=0.5, ls=0.1\n):\n    \"\"\"\n    Train one epoch.  Key defaults aligned with mogen:\n      fmp  = 0.5   (full_mask_prob -- forces text dependence)\n      ls   = 0.1   (label smoothing)\n    \"\"\"\n    if not res_only:\n        base_m.train()\n    else:\n        base_m.eval()\n\n    if res_m is not None:\n        res_m.train()\n\n    tl = ta = rl = ra = 0.0\n    nbat = nres = 0\n\n    pbar = tqdm(loader, desc=f\"Ep {ep}\")\n\n    for bi, (txts, toks, lens) in enumerate(pbar):\n        toks = toks.to(dev, non_blocking=True)\n        lens = lens.to(dev, non_blocking=True)\n\n        if len(toks.shape) == 3:\n            bt = toks[:, 0]\n            layers = [toks[:, i] for i in range(toks.shape[1])]\n        else:\n            bt = toks\n            layers = [toks]\n\n        accum_step = (bi + 1) % accum != 0\n\n        # ---- Base model ----\n        blv = ba = 0.0\n        if not res_only:\n            with autocast('cuda', enabled=amp):\n                bl, _, ba = base_m(bt, txts, lens, full_mask_prob=fmp, label_smoothing=ls)\n                bl = bl / accum\n\n            if amp and base_scaler:\n                base_scaler.scale(bl).backward()\n                if not accum_step:\n                    base_scaler.unscale_(base_opt)\n                    torch.nn.utils.clip_grad_norm_(base_m.params_no_clip(), 1.0)\n                    base_scaler.step(base_opt)\n                    base_scaler.update()\n                    base_opt.zero_grad()\n            else:\n                bl.backward()\n                if not accum_step:\n                    torch.nn.utils.clip_grad_norm_(base_m.params_no_clip(), 1.0)\n                    base_opt.step()\n                    base_opt.zero_grad()\n\n            blv = bl.item() * accum\n            tl += blv\n            ta += ba\n            nbat += 1\n\n        # ---- Residual model ----\n        rlv = batch_ra = 0.0\n        if train_res and res_m and vq and random.random() < res_p:\n            nq = len(layers)\n            li = random.randint(1, nq - 1)\n            prev = layers[:li]\n            targ = layers[li]\n\n            with autocast('cuda', enabled=amp):\n                batch_rl, _, batch_ra = res_m(prev, targ, li, txts, lens, vq)\n                batch_rl = batch_rl / accum\n\n            if amp and res_scaler:\n                res_scaler.scale(batch_rl).backward()\n                if not accum_step:\n                    res_scaler.unscale_(res_opt)\n                    torch.nn.utils.clip_grad_norm_(res_m.params_no_clip(), 1.0)\n                    res_scaler.step(res_opt)\n                    res_scaler.update()\n                    res_opt.zero_grad()\n            else:\n                batch_rl.backward()\n                if not accum_step:\n                    torch.nn.utils.clip_grad_norm_(res_m.params_no_clip(), 1.0)\n                    res_opt.step()\n                    res_opt.zero_grad()\n\n            rlv = batch_rl.item() * accum\n            rl += rlv\n            ra += batch_ra\n            nres += 1\n\n        # Progress bar\n        post = {}\n        if not res_only:\n            post['l'] = f'{blv:.4f}'\n            post['a'] = f'{ba:.4f}'\n        if rlv > 0:\n            post['rl'] = f'{rlv:.4f}'\n            post['ra'] = f'{batch_ra:.4f}'\n        if amp:\n            post['amp'] = 'on'\n        pbar.set_postfix(post)\n\n    # Step schedulers (after epoch, like mogen)\n    if base_sch: base_sch.step()\n    if res_sch:  res_sch.step()\n\n    # Epoch averages\n    m = {}\n    if nbat > 0:\n        m['l'] = tl / nbat\n        m['a'] = ta / nbat\n    if nres > 0:\n        m['rl'] = rl / nres\n        m['ra'] = ra / nres\n\n    # W&B logging\n    # if wandb.run:\n    #     wlog = {'epoch': ep}\n    #     if nbat > 0:\n    #         wlog.update({\n    #             'train/loss': m['l'],\n    #             'train/acc': m['a'],\n    #             'train/lr': base_opt.param_groups[0]['lr'],\n    #         })\n    #     if nres > 0:\n    #         wlog.update({\n    #             'train/res_loss': m['rl'],\n    #             'train/res_acc': m['ra'],\n    #         })\n    #     wandb.log(wlog)\n\n    return m\n\n\n@torch.no_grad()\ndef validate(model, loader, dev, ep, writer=None):\n    \"\"\"Validate base model (token-level CE + accuracy).\"\"\"\n    model.eval()\n\n    tl = ta = 0.0\n    nbat = 0\n\n    for bi, (txts, toks, lens) in enumerate(tqdm(loader, desc=\"Val\")):\n        toks = toks.to(dev)\n        lens = lens.to(dev)\n\n        bt = toks[:, 0] if len(toks.shape) == 3 else toks\n\n        l, _, a = model(bt, txts, lens)\n        tl += l.item()\n        ta += a\n        nbat += 1\n\n    out = {\n        'l': tl / nbat if nbat > 0 else 0,\n        'a': ta / nbat if nbat > 0 else 0,\n    }\n\n    # if wandb.run:\n    #     wandb.log({'epoch': ep, 'val/loss': out['l'], 'val/acc': out['a']})\n\n    return out\n\n\n@torch.no_grad()\ndef validate_residual(res_m, base_m, loader, dev, ep, vq):\n    \"\"\"\n    Validate residual transformer (token-level CE + accuracy).\n    Randomly samples a residual layer per batch like mogen's training loop.\n    \"\"\"\n    res_m.eval()\n    base_m.eval()\n\n    rl = ra = 0.0\n    nres = 0\n\n    for bi, (txts, toks, lens) in enumerate(tqdm(loader, desc=\"Val-Res\")):\n        toks = toks.to(dev)\n        lens = lens.to(dev)\n\n        if len(toks.shape) == 3:\n            layers = [toks[:, i] for i in range(toks.shape[1])]\n        else:\n            continue\n\n        nq = len(layers)\n        li = random.randint(1, nq - 1)\n        prev = layers[:li]\n        targ = layers[li]\n\n        loss, _, acc = res_m(prev, targ, li, txts, lens, vq)\n        rl += loss.item()\n        ra += acc\n        nres += 1\n\n    out = {\n        'rl': rl / nres if nres > 0 else 0,\n        'ra': ra / nres if nres > 0 else 0,\n    }\n\n    # if wandb.run:\n    #     wandb.log({'epoch': ep, 'val/res_loss': out['rl'], 'val/res_acc': out['ra']})\n\n    return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:40.088817Z","iopub.execute_input":"2026-04-28T14:25:40.089147Z","iopub.status.idle":"2026-04-28T14:25:40.119308Z","shell.execute_reply.started":"2026-04-28T14:25:40.089125Z","shell.execute_reply":"2026-04-28T14:25:40.118676Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Initialization and Training Execution\n\n**Model Setup**: Instantiates both transformer architectures using the configuration parameters defined earlier:\n- **MaskTransformer**: Generates the base layer (layer 0) through iterative masked token prediction\n- **ResidualTransformer**: Generates refinement layers (layers 1-5) autoregressively\n\nBoth models share the same architectural hyperparameters but differ in their generation strategies.\n\n**Data Loading and Preprocessing**:\n1. Reads the training CSV containing pre-tokenized motion sequences\n2. Parses space-separated token strings into numpy arrays\n3. Stacks all 6 RVQ layers (base + 5 residual) into a single tensor per sample\n4. Handles variable-length sequences by truncating to minimum layer length\n5. Creates a 90/10 train/validation split\n\n**Training Configuration**:\n- Uses AdamW optimizer for stable training with weight decay\n- Batch processing with custom collate function to handle text + tokens\n- Trains only the base model in this cell (residual model training shown later)\n\n**Training Loop**: Iterates through epochs, alternating between training and validation phases. Tracks loss and accuracy metrics for both phases.\n\n**Visualization**: Plots training curves showing loss and accuracy progression over epochs, saved to the output directory for monitoring convergence and detecting overfitting.","metadata":{}},{"cell_type":"code","source":"class VectorQuantizer(nn.Module):\n    def __init__(self, num_embeddings=512, embedding_dim=256):\n        super().__init__()\n        self.embedding_dim = embedding_dim\n        self.num_embeddings = num_embeddings\n        self.register_buffer('embedding', torch.randn(embedding_dim, num_embeddings))\n\n\nclass ResidualVectorQuantizer(nn.Module):\n    def __init__(self, num_quantizers=6, num_embeddings=512, embedding_dim=256):\n        super().__init__()\n        self.num_quantizers = num_quantizers\n        self.quantizers = nn.ModuleList([\n            VectorQuantizer(num_embeddings, embedding_dim)\n            for _ in range(num_quantizers)\n        ])\n\n    def quantize_from_tokens(self, tokens):\n        codes = []\n        for v, toks in enumerate(tokens):\n            B, n = toks.shape\n            d = self.quantizers[v].embedding_dim\n            flat = toks.reshape(-1)\n            q = self.quantizers[v].embedding[:, flat].reshape(d, B, n).permute(1, 0, 2)\n            codes.append(q)\n        return sum(codes)\n\n\nclass RVQVAE(nn.Module):\n    def __init__(self, encoder, decoder, rvq, downsampling_ratio=4):\n        super().__init__()\n        self.encoder = encoder\n        self.decoder = decoder\n        self.rvq = rvq\n        self.num_quantizers = rvq.num_quantizers\n        self.downsampling_ratio = downsampling_ratio\n\n\ndef load_vae(path, device='cuda'):\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\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\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\n    # Read downsampling_ratio from checkpoint config (default 4)\n    ds_ratio = cfg.get('downsampling_ratio', 4)\n\n    vae = RVQVAE(None, None, rvq, downsampling_ratio=ds_ratio)\n    vae.to(device).eval()\n    print(f\"VAE loaded: {cfg.get('num_quantizers',6)} quantizers, \"\n          f\"{cfg.get('num_embeddings',512)} embeddings, \"\n          f\"downsampling_ratio={ds_ratio}\")\n    return vae\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:40.120195Z","iopub.execute_input":"2026-04-28T14:25:40.120497Z","iopub.status.idle":"2026-04-28T14:25:40.136446Z","shell.execute_reply.started":"2026-04-28T14:25:40.120466Z","shell.execute_reply":"2026-04-28T14:25:40.135616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n\nset_seed(42)\n# init_wandb(\n#     config={**VAE_CONFIG, **TRANSFORMER_CONFIG},\n#     run_name='motion-s-aligned'\n# )\n\n# ===== DATA =====\ndf = pd.read_csv(CSV_PATH)\nvq_model = load_vae(\n    VAE_PATH, device\n)\n\ndef parse(s):\n    if pd.isna(s): return np.array([])\n    return np.array([int(x) for x in str(s).split() if x.strip()], dtype=np.int32)\n\ntoken_cols = [\n    'base_tokens', 'residual_1', 'residual_2',\n    'residual_3', 'residual_4', 'residual_5'\n]\nfor c in token_cols:\n    df[c] = df[c].apply(parse)\n\n# Use base_tokens length as canonical (like mogen uses base layer length)\ntoken_data = {}\nfor _, r in df.iterrows():\n    sid = str(r['id'])\n    base = r['base_tokens']\n    if len(base) == 0:\n        continue\n    canon_len = len(base)\n\n    layers = []\n    for c in token_cols:\n        t = r[c]\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\n    tokens = np.stack(layers)  # (num_quantizers, canon_len)\n    token_data[sid] = {\n        'tokens': tokens,\n        'sentence': r['sentence'],\n        'gloss': r['gloss'],\n    }\n\n# Seeded random split (like mogen)\nkeys = list(token_data.keys())\nrandom.shuffle(keys)  # uses seed from set_seed(42)\n\nsplit = int(len(keys) * 0.9)\ntrain_keys, val_keys = keys[:split], keys[split:]\n\ntrain_d = {k: token_data[k] for k in train_keys}\nval_d   = {k: token_data[k] for k in val_keys}\n\nC = TRANSFORMER_CONFIG\ntrain_ds = TokenDataset(train_d, C['text_source'], C['max_token_len'], C['min_token_len'])\nval_ds   = TokenDataset(val_d,   C['text_source'], C['max_token_len'], C['min_token_len'])\n\n\ndef collate(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\n\ntrain_loader = DataLoader(\n    train_ds, batch_size=C['batch_size'], shuffle=True,\n    collate_fn=collate, pin_memory=True, drop_last=True\n)\nval_loader = DataLoader(\n    val_ds, batch_size=C['batch_size'], shuffle=False,\n    collate_fn=collate, pin_memory=True\n)\n\n# ===== BUILD MODELS =====\nmax_seq_len = C['max_token_len'] + 77  # motion + CLIP text tokens (like mogen)\n\ntransformer = MaskTransformer(\n    num_tokens=VAE_CONFIG['num_embeddings'],\n    code_dim=VAE_CONFIG['latent_dim'],\n    latent_dim=C['latent_dim'],\n    ff_size=C['ff_size'],\n    num_layers=C['num_layers'],\n    num_heads=C['num_heads'],\n    dropout=C['dropout'],\n    clip_dim=512,\n    clip_version=\"ViT-B/32\",\n    cond_drop_prob=C['cond_drop_prob'],\n    device=device,\n    max_seq_len=max_seq_len,\n).to(device)\n\n# Residual stays on CPU until needed (saves ~2GB VRAM during Phase 1)\nres_model = ResidualTransformer(\n    num_tokens=VAE_CONFIG['num_embeddings'],\n    code_dim=VAE_CONFIG['latent_dim'],\n    latent_dim=C['latent_dim'],\n    ff_size=C['ff_size'],\n    num_layers=C['num_layers'],\n    num_heads=C['num_heads'],\n    dropout=C['dropout'],\n    clip_dim=512,\n    clip_version=\"ViT-B/32\",\n    cond_drop_prob=C['cond_drop_prob'],\n    device=\"cpu\",\n    max_seq_len=max_seq_len,\n    num_quantizers=VAE_CONFIG['num_quantizers'],\n)\n\n# ===== OPTIMIZERS + LR SCHEDULES (warmup + cosine, like mogen) =====\nbase_opt = torch.optim.AdamW(\n    transformer.params_no_clip(),\n    lr=C['lr'],\n    weight_decay=C['weight_decay'],\n)\nbase_sch = torch.optim.lr_scheduler.LambdaLR(\n    base_opt, make_lr_lambda(C['warmup_epochs'], C['epochs'])\n)\n\nuse_amp = True\nbase_scaler = torch.amp.GradScaler(\"cuda\") if use_amp else None\nres_scaler = torch.amp.GradScaler(\"cuda\") if use_amp else None\n\n# Residual optimizer created lazily when needed\nres_opt = None\nres_sch = None\n\n# ===== PARAM COUNTS =====\ntotal_p = sum(p.numel() for p in transformer.parameters())\ntrain_p = sum(p.numel() for p in transformer.params_no_clip())\nprint(f\"MaskTransformer: {total_p:,} total, {train_p:,} trainable\")\n\nres_total_p = sum(p.numel() for p in res_model.parameters())\nres_train_p = sum(p.numel() for p in res_model.params_no_clip())\nprint(f\"ResidualTransformer: {res_total_p:,} total, {res_train_p:,} trainable\")\n\n# ===== SINGLE TRAINING LOOP (mogen-style staged) =====\nprint(\"\\n\" + \"=\"*60)\nprint(\"TRAINING: single loop with staged residual (mogen-style)\")\nprint(f\"  Effective batch: {C['batch_size']} x {C['grad_accum']} = {C['batch_size']*C['grad_accum']}\")\nprint(f\"  Epochs: {C['epochs']}, Warmup: {C['warmup_epochs']}\")\nprint(f\"  full_mask_prob={C['full_mask_prob']}, cond_drop_prob={C['cond_drop_prob']}\")\nprint(f\"  Residual starts at epoch {C['residual_start_epoch']}\")\nprint(\"=\"*60 + \"\\n\")\n\nhistory = {\n    'ep': [], 'tl': [], 'ta': [], 'vl': [], 'va': [],\n    'trl': [], 'tra': [], 'vrl': [], 'vra': [],\n}\n\nbest_val_loss = float('inf')\nbest_res_loss = float('inf')\n\nfor ep in range(1, C['epochs'] + 1):\n\n    # Determine if we train residual this epoch\n    train_res = (ep >= C['residual_start_epoch'])\n\n    # Lazily move res_model to GPU and create optimizer on first residual epoch\n    if train_res and res_opt is None:\n        print(f\"\\n*** Epoch {ep}: activating ResidualTransformer on GPU ***\")\n        torch.cuda.empty_cache()\n        res_model = res_model.to(device)\n        res_model.tenc = transformer.tenc  # share frozen CLIP encoder\n        res_opt = torch.optim.AdamW(\n            res_model.params_no_clip(),\n            lr=C['res_lr'],\n            weight_decay=C['res_weight_decay'],\n        )\n        res_sch = torch.optim.lr_scheduler.LambdaLR(\n            res_opt, make_lr_lambda(C['warmup_epochs'], C['epochs'] - C['residual_start_epoch'])\n        )\n\n    ## \n\n    # ---- Train ----\n    m = train_epoch(\n        base_m=transformer,\n        loader=train_loader,\n        base_opt=base_opt,\n        base_sch=base_sch,\n        dev=device,\n        ep=ep,\n        res_m=res_model if train_res else None,\n        res_opt=res_opt,\n        res_sch=res_sch,\n        vq=vq_model if train_res else None,\n        train_res=train_res,\n        res_p=C['res_prob'],\n        amp=use_amp,\n        base_scaler=base_scaler,\n        res_scaler=res_scaler,\n        accum=C['grad_accum'],\n        res_only=False,             # always train base too (mogen default)\n        fmp=C['full_mask_prob'],\n        ls=C['label_smoothing'],\n    )\n\n    # ---- Validate base ----\n    vm = validate(transformer, val_loader, device, ep)\n\n    # ---- Validate residual (if active) ----\n    vrm = {}\n    if train_res and res_opt is not None:\n        vrm = validate_residual(res_model, transformer, val_loader, device, ep, vq_model)\n\n    # ---- Record history ----\n    history['ep'].append(ep)\n    history['tl'].append(m.get('l', 0.))\n    history['ta'].append(m.get('a', 0.))\n    history['vl'].append(vm.get('l', 0.))\n    history['va'].append(vm.get('a', 0.))\n    history['trl'].append(m.get('rl', 0.))\n    history['tra'].append(m.get('ra', 0.))\n    history['vrl'].append(vrm.get('rl', 0.))\n    history['vra'].append(vrm.get('ra', 0.))\n\n    # ---- Print ----\n    line = f\"{ep:3d}/{C['epochs']} | tr {m.get('l',0.):.4f} {m.get('a',0.):.3f}\"\n    line += f\" | val {vm['l']:.4f} {vm['a']:.3f}\"\n    if train_res:\n        line += f\" | res tr {m.get('rl',0.):.4f} {m.get('ra',0.):.3f}\"\n        line += f\" val {vrm.get('rl',0.):.4f} {vrm.get('ra',0.):.3f}\"\n    print(line)\n\n    # ---- Save best base ----\n    if vm['l'] < best_val_loss:\n        best_val_loss = vm['l']\n        ckpt = {\n            'epoch': ep,\n            'model_state_dict': transformer.state_dict(),\n            'optimizer_state_dict': base_opt.state_dict(),\n            'metrics': {'train': m, 'val': vm},\n        }\n        if train_res and res_opt is not None:\n            ckpt['residual_model_state_dict'] = res_model.state_dict()\n            ckpt['residual_optimizer_state_dict'] = res_opt.state_dict()\n        torch.save(ckpt, str(OUTPUT_ROOT / \"best_model.pth\"))\n        print(f\"  -> best base model saved (val_loss={best_val_loss:.4f})\")\n\n    # ---- Save best residual ----\n    if train_res and vrm.get('rl', float('inf')) < best_res_loss:\n        best_res_loss = vrm['rl']\n        ckpt = {\n            'epoch': ep,\n            'model_state_dict': transformer.state_dict(),\n            'residual_model_state_dict': res_model.state_dict(),\n            'residual_optimizer_state_dict': res_opt.state_dict(),\n            'metrics': {'train': m, 'val': vm, 'val_res': vrm},\n        }\n        torch.save(ckpt, str(OUTPUT_ROOT / \"best_residual_model.pth\"))\n        print(f\"  -> best residual model saved (val_res_loss={best_res_loss:.4f})\")\n\n    # ---- Periodic checkpoint ----\n    if ep % C['save_every'] == 0:\n        ckpt = {\n            'epoch': ep,\n            'model_state_dict': transformer.state_dict(),\n            'optimizer_state_dict': base_opt.state_dict(),\n            'metrics': {'train': m, 'val': vm},\n        }\n        if train_res and res_opt is not None:\n            ckpt['residual_model_state_dict'] = res_model.state_dict()\n            ckpt['residual_optimizer_state_dict'] = res_opt.state_dict()\n        torch.save(ckpt, str(OUTPUT_ROOT / f\"checkpoint_ep{ep}.pth\"))\n\n# ===== FINAL SAVE =====\nckpt = {\n    'epoch': C['epochs'],\n    'model_state_dict': transformer.state_dict(),\n    'optimizer_state_dict': base_opt.state_dict(),\n}\nif res_opt is not None:\n    ckpt['residual_model_state_dict'] = res_model.state_dict()\n    ckpt['residual_optimizer_state_dict'] = res_opt.state_dict()\ntorch.save(ckpt, str(OUTPUT_ROOT / \"final_model.pth\"))\n\n# ===== PLOTS =====\nfig, axes = plt.subplots(2, 2, figsize=(12, 8))\n\n# Base loss\naxes[0, 0].plot(history['ep'], history['tl'], label='train')\naxes[0, 0].plot(history['ep'], history['vl'], label='val')\naxes[0, 0].set_title('Base Loss')\naxes[0, 0].legend(); axes[0, 0].grid(alpha=0.3)\n\n# Base accuracy\naxes[0, 1].plot(history['ep'], history['ta'], label='train')\naxes[0, 1].plot(history['ep'], history['va'], label='val')\naxes[0, 1].set_title('Base Accuracy')\naxes[0, 1].legend(); axes[0, 1].grid(alpha=0.3)\n\n# Residual loss\naxes[1, 0].plot(history['ep'], history['trl'], label='train')\naxes[1, 0].plot(history['ep'], history['vrl'], label='val')\naxes[1, 0].set_title('Residual Loss')\naxes[1, 0].legend(); axes[1, 0].grid(alpha=0.3)\n\n# Residual accuracy\naxes[1, 1].plot(history['ep'], history['tra'], label='train')\naxes[1, 1].plot(history['ep'], history['vra'], label='val')\naxes[1, 1].set_title('Residual Accuracy')\naxes[1, 1].legend(); axes[1, 1].grid(alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(OUTPUT_ROOT / \"training_curves.png\", dpi=150)\nplt.show()\n\n# import wandb\n# wandb.finish()\nprint(\"\\nAll training complete!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T14:25:40.137420Z","iopub.execute_input":"2026-04-28T14:25:40.137709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.distributions.categorical import Categorical\nfrom typing import List, Optional, Tuple, Union\n\nclass LengthEstimator(nn.Module):\n    \"\"\"\n    Neural network to predict motion sequence length from text embeddings.\n    \n    Architecture (following MoMask):\n        CLIP embedding (512) → MLP → Length bin logits (num_bins)\n    \n    The model predicts a distribution over discrete length bins.\n    During inference, you can:\n        - Sample from the distribution (stochastic, more diverse)\n        - Take argmax (deterministic)\n        - Use top-k sampling (balanced)\n    \n    Args:\n        clip_dim: Dimension of CLIP text embeddings (512 for ViT-B/32)\n        num_bins: Number of discrete length bins (default 50)\n        hidden_dim: Hidden layer dimension (default 512)\n        min_tokens: Minimum token length (bin 0 corresponds to this)\n        max_tokens: Maximum token length (last bin corresponds to this)\n        dropout: Dropout rate\n    \"\"\"\n    \n    def __init__(\n        self,\n        clip_dim: int = 512,\n        num_bins: int = 50,\n        hidden_dim: int = 512,\n        min_tokens: int = 10,\n        max_tokens: int = 200,\n        dropout: float = 0.2\n    ):\n        super().__init__()\n        \n        self.clip_dim = clip_dim\n        self.num_bins = num_bins\n        self.min_tokens = min_tokens\n        self.max_tokens = max_tokens\n        self.hidden_dim = hidden_dim\n        \n        # Calculate bin size\n        self.bin_size = (max_tokens - min_tokens) / (num_bins - 1)\n        \n        # MLP following MoMask architecture\n        self.net = nn.Sequential(\n            nn.Linear(clip_dim, hidden_dim),\n            nn.LayerNorm(hidden_dim),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Dropout(dropout),\n            \n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.LayerNorm(hidden_dim // 2),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Dropout(dropout),\n            \n            nn.Linear(hidden_dim // 2, hidden_dim // 4),\n            nn.LayerNorm(hidden_dim // 4),\n            nn.LeakyReLU(0.2, inplace=True),\n            \n            nn.Linear(hidden_dim // 4, num_bins)\n        )\n        \n        # Initialize weights\n        self.apply(self._init_weights)\n    \n    def _init_weights(self, module):\n        \"\"\"Initialize weights following MoMask.\"\"\"\n        if isinstance(module, nn.Linear):\n            nn.init.xavier_uniform_(module.weight)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        elif isinstance(module, nn.LayerNorm):\n            nn.init.ones_(module.weight)\n            nn.init.zeros_(module.bias)\n    \n    def forward(self, clip_embedding: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Forward pass to get length bin logits.\n        \n        Args:\n            clip_embedding: (B, clip_dim) CLIP text embeddings\n            \n        Returns:\n            logits: (B, num_bins) unnormalized log probabilities for each bin\n        \"\"\"\n        return self.net(clip_embedding)\n    \n    def tokens_to_bin(self, token_lengths: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Convert continuous token lengths to discrete bins.\n        \n        Args:\n            token_lengths: (B,) token lengths\n            \n        Returns:\n            bins: (B,) bin indices (0 to num_bins-1)\n        \"\"\"\n        # Linear mapping from [min_tokens, max_tokens] to [0, num_bins-1]\n        bins = ((token_lengths - self.min_tokens) / self.bin_size).round().long()\n        bins = bins.clamp(0, self.num_bins - 1)\n        return bins\n    \n    def bin_to_tokens(self, bins: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Convert discrete bins back to token lengths.\n        \n        Args:\n            bins: (B,) bin indices\n            \n        Returns:\n            token_lengths: (B,) token lengths\n        \"\"\"\n        token_lengths = self.min_tokens + bins.float() * self.bin_size\n        return token_lengths.round().long()\n    \n    @torch.no_grad()\n    def predict(\n        self,\n        clip_embedding: torch.Tensor,\n        mode: str = 'sample',\n        temperature: float = 1.0\n    ) -> torch.Tensor:\n        \"\"\"\n        Predict token lengths from CLIP embeddings.\n        \n        Args:\n            clip_embedding: (B, clip_dim) CLIP text embeddings\n            mode: Prediction mode\n                - 'sample': Sample from distribution (stochastic)\n                - 'argmax': Take most likely bin (deterministic)\n                - 'mean': Weighted average of bins (expected value)\n            temperature: Sampling temperature (only for mode='sample')\n            \n        Returns:\n            token_lengths: (B,) predicted token lengths\n        \"\"\"\n        self.eval()\n        logits = self.forward(clip_embedding)\n        \n        if mode == 'argmax':\n            bins = logits.argmax(dim=-1)\n        elif mode == 'sample':\n            probs = F.softmax(logits / temperature, dim=-1)\n            bins = Categorical(probs).sample()\n        elif mode == 'mean':\n            probs = F.softmax(logits, dim=-1)\n            bin_indices = torch.arange(self.num_bins, device=logits.device, dtype=torch.float)\n            bins = (probs * bin_indices).sum(dim=-1).round().long()\n        else:\n            raise ValueError(f\"Unknown mode: {mode}. Use 'sample', 'argmax', or 'mean'.\")\n        \n        return self.bin_to_tokens(bins)\n    \n    @torch.no_grad()\n    def estimate_lengths(\n        self,\n        texts: List[str],\n        clip_model: nn.Module,\n        mode: str = 'sample',\n        temperature: float = 1.0\n    ) -> List[int]:\n        \"\"\"\n        Estimate motion lengths from text strings.\n        \n        Convenience method that handles CLIP encoding internally.\n        \n        Args:\n            texts: List of text prompts\n            clip_model: CLIP text encoder with encode_text method\n            mode: Prediction mode ('sample', 'argmax', 'mean')\n            temperature: Sampling temperature\n            \n        Returns:\n            List of estimated token lengths\n        \"\"\"\n        self.eval()\n        \n        # Get CLIP embeddings\n        clip_emb = clip_model(texts)  # (B, clip_dim)\n        \n        # Predict lengths\n        token_lengths = self.predict(clip_emb, mode=mode, temperature=temperature)\n        \n        return token_lengths.tolist()\n\n\ndef load_length_estimator(\n    checkpoint_path: str,\n    device: Union[str, torch.device] = 'cuda'\n) -> Tuple[LengthEstimator, dict]:\n    \"\"\"\n    Load a trained LengthEstimator from checkpoint.\n    \n    Args:\n        checkpoint_path: Path to checkpoint file\n        device: Device to load model on\n        \n    Returns:\n        Tuple of (model, config_dict)\n    \"\"\"\n    checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)\n    \n    config = checkpoint.get('config', {})\n    \n    model = LengthEstimator(\n        clip_dim=config.get('clip_dim', 512),\n        num_bins=config.get('num_bins', 50),\n        hidden_dim=config.get('hidden_dim', 512),\n        min_tokens=config.get('min_tokens', 10),\n        max_tokens=config.get('max_tokens', 200),\n        dropout=0.0  # No dropout during inference\n    )\n    \n    # Load weights\n    if 'model_state_dict' in checkpoint:\n        model.load_state_dict(checkpoint['model_state_dict'])\n    elif 'estimator' in checkpoint:\n        # MoMask format compatibility\n        model.load_state_dict(checkpoint['estimator'])\n    else:\n        model.load_state_dict(checkpoint)\n    \n    model.to(device)\n    model.eval()\n    \n    return model, config\n\n\nclass LengthEstimatorWithCLIP(nn.Module):\n    \"\"\"\n    Combined module with CLIP encoder + Length Estimator.\n    \n    Convenience wrapper that includes the CLIP model for end-to-end inference.\n    \"\"\"\n    \n    def __init__(\n        self,\n        length_estimator: LengthEstimator,\n        clip_model: nn.Module\n    ):\n        super().__init__()\n        self.length_estimator = length_estimator\n        self.clip_model = clip_model\n        \n        # Freeze CLIP\n        for param in self.clip_model.parameters():\n            param.requires_grad = False\n    \n    def forward(self, texts: List[str]) -> torch.Tensor:\n        \"\"\"Get length logits from text strings.\"\"\"\n        clip_emb = self.clip_model.encode_text(texts)\n        return self.length_estimator(clip_emb)\n    \n    @torch.no_grad()\n    def estimate(\n        self,\n        texts: List[str],\n        mode: str = 'sample',\n        temperature: float = 1.0\n    ) -> List[int]:\n        \"\"\"Estimate motion lengths from text.\"\"\"\n        self.eval()\n        clip_emb = self.clip_model.encode_text(texts)\n        token_lengths = self.length_estimator.predict(clip_emb, mode=mode, temperature=temperature)\n        return token_lengths.tolist()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"length_estimator, le_config = load_length_estimator(\"/kaggle/input/models/antonygithinji/motion-s-length-estimator/pytorch/default/1/length_estimator_best.pth\", device)\n\n# Your CLIP encoder is at transformer.tenc (not .text_encoder)\nclip_model = transformer.tenc\n\n# Define the function your generation loop uses\ndef estimate_length(text):\n    lengths = length_estimator.estimate_lengths(\n        [text],          # list of strings\n        clip_model,      # your transformer's CLIPTextEncoder\n        mode='mean',\n        temperature=1.0\n    )\n    return lengths[0]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fix device mismatch for text encoders\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Clear caches and update device for both models' text encoders\nif hasattr(transformer, 'tenc'):\n    transformer.tenc.device = device\n    transformer.tenc._cache.clear()\n    \nif hasattr(res_model, 'tenc'):\n    res_model.tenc.device = device\n    res_model.tenc._cache.clear()\n\nprint(f\"Text encoder caches cleared, device set to: {device}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Define the to_str function\ndef to_str(tokens):\n    \"\"\"Convert tensor tokens to space-separated string\"\"\"\n    if isinstance(tokens, torch.Tensor):\n        tokens = tokens.cpu().numpy()\n    return ' '.join(map(str, tokens.astype(int).flatten().tolist()))\n\n\n# Define estimate_length function\ndef estimate_length(text):\n    lengths = length_estimator.estimate_lengths(\n        [text],          # list of strings\n        clip_model,      # your transformer's CLIPTextEncoder\n        mode='mean',\n        temperature=1.0\n    )\n    return lengths[0]\n\n# Fix circular references\ndef fix_circular_references(module, visited=None):\n    \"\"\"Recursively find and break circular references\"\"\"\n    if visited is None:\n        visited = set()\n    \n    module_id = id(module)\n    if module_id in visited:\n        return\n    visited.add(module_id)\n    \n    # Check all attributes\n    for name in list(module._modules.keys()):\n        child = module._modules[name]\n        if child is not None:\n            # If a child module references itself, break the cycle\n            if hasattr(child, '_modules'):\n                for child_name in list(child._modules.keys()):\n                    if child._modules[child_name] is child:\n                        print(f\"Breaking circular reference: {child.__class__.__name__}.{child_name} -> {child.__class__.__name__}\")\n                        child._modules[child_name] = None\n            \n            # Recursively check children\n            fix_circular_references(child, visited)\n\nprint(\"Fixing circular references in transformer...\")\nfix_circular_references(transformer)\n\nprint(\"Fixing circular references in res_model...\")\nfix_circular_references(res_model)\n\n# Set training=False recursively to avoid calling .eval()\ndef set_training_false_recursive(module, visited=None):\n    if visited is None:\n        visited = set()\n    \n    module_id = id(module)\n    if module_id in visited:\n        return\n    visited.add(module_id)\n    \n    module.training = False\n    for child in module._modules.values():\n        if child is not None:\n            set_training_false_recursive(child, visited)\n\nset_training_false_recursive(transformer)\nset_training_false_recursive(res_model)\nset_training_false_recursive(vq_model)\nset_training_false_recursive(length_estimator)\n\nprint(\"Models set to eval mode. Starting generation...\")\n\n# Generate predictions\nresults = []\nfor _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Generating\"):\n    sid = row['id']\n    text = f\"{row['sentence']} {row['gloss']}\"\n    \n    # Estimate length\n    tlen = estimate_length(text)\n    mlens = torch.tensor([tlen], device=device)\n    \n    with torch.no_grad():\n        # Generate base tokens\n        base = transformer.generate([text], mlens, timesteps=10, cond_scale=4.0, temperature=1.0)\n        base = torch.clamp(base, 0, 511)\n        \n        # Generate residual layers\n        layers = [base]\n        for li in range(1, 6):\n            res = res_model.generate_layer(layers, li, [text], mlens, vq_model, cs=4.0, temp=1.0, th=0.9)\n            res = torch.clamp(res, 0, 511)\n            layers.append(res)\n        \n        # Convert to strings\n        row_data = {\n            'id': sid,\n            'base_tokens': to_str(base[0]),\n            'residual_1': to_str(layers[1][0]),\n            'residual_2': to_str(layers[2][0]),\n            'residual_3': to_str(layers[3][0]),\n            'residual_4': to_str(layers[4][0]),\n            'residual_5': to_str(layers[5][0])\n        }\n        results.append(row_data)\n\n# Save submission\nsubmission = pd.DataFrame(results)\nsubmission.to_csv('submission.csv', index=False)\nprint(f\"\\n✓ Generated {len(submission)} predictions\")\nprint(f\"✓ Saved to submission.csv\")\nsubmission.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}