{"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}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Install & Imports AND Load Data","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import AutoTokenizer, AutoModel\nfrom tqdm import tqdm\n\nDATA_PATH = \"/kaggle/input/motion-s-hierarchical-text-to-motion-generation-for-sign-language\"\n\ntrain_df = pd.read_csv(f\"{DATA_PATH}/train.csv\")\ntest_df  = pd.read_csv(f\"{DATA_PATH}/test.csv\")\n\ndef parse_tokens(x):\n    if pd.isna(x) or not isinstance(x, str):\n        return np.array([])\n    return np.array(list(map(int, x.split())))\n\nfor col in ['base_tokens','residual_1','residual_2','residual_3','residual_4','residual_5']:\n    train_df[col] = train_df[col].apply(parse_tokens)\n\ntrain_df[\"length\"] = train_df[\"base_tokens\"].apply(len)\n\n# Filter out bad rows (empty tokens, length out of range)\ntrain_df = train_df[train_df[\"length\"].between(40, 800)].reset_index(drop=True)\nprint(f\"Training samples: {len(train_df)}, mean length: {train_df['length'].mean():.1f}\")\n\nMAX_LEN    = 256\nVOCAB_SIZE = 512\nNUM_LAYERS = 6","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-26T13:57:47.137023Z","iopub.execute_input":"2026-02-26T13:57:47.137289Z","iopub.status.idle":"2026-02-26T13:58:04.474371Z","shell.execute_reply.started":"2026-02-26T13:57:47.137240Z","shell.execute_reply":"2026-02-26T13:58:04.473490Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset Class","metadata":{}},{"cell_type":"code","source":"class MotionDataset(Dataset):\n    def __init__(self, df, tokenizer, max_text_len=64, max_motion_len=256):\n        self.df = df.reset_index(drop=True)\n        self.tokenizer = tokenizer\n        self.max_text_len = max_text_len\n        self.max_motion_len = max_motion_len\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        # Combine gloss + sentence for richer signal\n        text = str(row[\"gloss\"]) + \" [SEP] \" + str(row[\"sentence\"])\n        \n        encoding = self.tokenizer(\n            text,\n            padding=\"max_length\",\n            truncation=True,\n            max_length=self.max_text_len,\n            return_tensors=\"pt\"\n        )\n        \n        token_arrays = [\n            row['base_tokens'], row['residual_1'], row['residual_2'],\n            row['residual_3'],  row['residual_4'], row['residual_5']\n        ]\n        \n        motion_len = min(len(token_arrays[0]), self.max_motion_len)\n        \n        padded = np.zeros((6, self.max_motion_len), dtype=np.int64)\n        for i, arr in enumerate(token_arrays):\n            arr_len = min(len(arr), self.max_motion_len)\n            padded[i, :arr_len] = arr[:arr_len]\n        \n        return {\n            \"input_ids\":      encoding[\"input_ids\"].squeeze(0),\n            \"attention_mask\": encoding[\"attention_mask\"].squeeze(0),\n            \"tokens\":         torch.tensor(padded, dtype=torch.long),\n            \"motion_len\":     motion_len,\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T13:58:57.933565Z","iopub.execute_input":"2026-02-26T13:58:57.934230Z","iopub.status.idle":"2026-02-26T13:58:57.941909Z","shell.execute_reply.started":"2026-02-26T13:58:57.934195Z","shell.execute_reply":"2026-02-26T13:58:57.941103Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Architecture","metadata":{}},{"cell_type":"code","source":"class ImprovedMotionModel(nn.Module):\n    def __init__(self, hidden=768, num_decoder_layers=4):\n        super().__init__()\n        \n        # Text encoder (BERT)\n        self.encoder = AutoModel.from_pretrained(\"bert-base-uncased\")\n        \n        # Project BERT hidden to our hidden size if needed\n        self.enc_proj = nn.Linear(768, hidden) if hidden != 768 else nn.Identity()\n        \n        # Layer-specific embedding: tells decoder which RVQ layer it's generating\n        self.layer_embed = nn.Embedding(NUM_LAYERS, hidden)\n        \n        # Shared token embedding (coarse tokens inform fine layers)\n        self.token_embed = nn.Embedding(VOCAB_SIZE + 1, hidden, padding_idx=VOCAB_SIZE)\n        \n        # Deeper decoder\n        decoder_layer = nn.TransformerDecoderLayer(\n            d_model=hidden,\n            nhead=8,\n            dim_feedforward=2048,\n            dropout=0.1,\n            batch_first=True,\n            norm_first=True  # Pre-norm for more stable training\n        )\n        self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_decoder_layers)\n        \n        # Per-layer output heads (slightly different projection per layer)\n        self.output_heads = nn.ModuleList([\n            nn.Sequential(\n                nn.Linear(hidden, hidden),\n                nn.GELU(),\n                nn.Linear(hidden, VOCAB_SIZE)\n            )\n            for _ in range(NUM_LAYERS)\n        ])\n        \n        # Length prediction head (from [CLS] token)\n        self.length_head = nn.Linear(hidden, 1)\n        \n    def encode_text(self, input_ids, attention_mask):\n        enc_out = self.encoder(\n            input_ids=input_ids,\n            attention_mask=attention_mask\n        ).last_hidden_state  # (B, T, 768)\n        return self.enc_proj(enc_out)\n    \n    def forward(self, input_ids, attention_mask, tgt_tokens, layer_idx):\n        \"\"\"\n        tgt_tokens: (B, S) - token sequence for this layer\n        layer_idx:  int    - which RVQ layer (0-5)\n        \"\"\"\n        B, S = tgt_tokens.shape\n        \n        memory = self.encode_text(input_ids, attention_mask)\n        \n        # Token embeddings\n        tgt_emb = self.token_embed(tgt_tokens)  # (B, S, H)\n        \n        # Add layer identity to every position\n        layer_id = torch.full((B,), layer_idx, dtype=torch.long, device=tgt_tokens.device)\n        layer_vec = self.layer_embed(layer_id).unsqueeze(1)  # (B, 1, H)\n        tgt_emb = tgt_emb + layer_vec\n        \n        # Causal mask\n        tgt_mask = torch.triu(\n            torch.ones(S, S, device=tgt_tokens.device), diagonal=1\n        ).bool()\n        \n        # Decode\n        dec_out = self.decoder(tgt=tgt_emb, memory=memory, tgt_mask=tgt_mask)\n        \n        logits = self.output_heads[layer_idx](dec_out)  # (B, S, VOCAB_SIZE)\n        \n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T13:59:03.728209Z","iopub.execute_input":"2026-02-26T13:59:03.728912Z","iopub.status.idle":"2026-02-26T13:59:03.738921Z","shell.execute_reply.started":"2026-02-26T13:59:03.728869Z","shell.execute_reply":"2026-02-26T13:59:03.737972Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ntokenizer = AutoTokenizer.from_pretrained(\"bert-base-uncased\")\n\ndataset = MotionDataset(train_df, tokenizer, max_text_len=64, max_motion_len=MAX_LEN)\nloader  = DataLoader(dataset, batch_size=16, shuffle=True, num_workers=2, pin_memory=True)\n\nmodel = ImprovedMotionModel(hidden=768, num_decoder_layers=4).to(device)\n\nNUM_EPOCHS = 5\n\n# Two param groups from the start: BERT (low lr) + rest (high lr)\noptimizer = torch.optim.AdamW([\n    {\"params\": model.encoder.parameters(),     \"lr\": 2e-5,  \"weight_decay\": 1e-4},\n    {\"params\": [p for name, p in model.named_parameters() if not name.startswith(\"encoder\")],\n               \"lr\": 3e-4, \"weight_decay\": 1e-4},\n], lr=3e-4)\n\n# Freeze BERT for epoch 0 by zeroing its lr\ndef set_bert_lr(optimizer, lr):\n    optimizer.param_groups[0][\"lr\"] = lr\n\n# Simple cosine scheduler — no issues with dynamic param groups\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=NUM_EPOCHS * len(loader), eta_min=1e-6\n)\n\ncriterion = nn.CrossEntropyLoss(ignore_index=0)\n\n# ---- Training ----\nfor epoch in range(NUM_EPOCHS):\n\n    # Epoch 0: freeze BERT (lr=0), epoch 1+: unfreeze\n    if epoch == 0:\n        set_bert_lr(optimizer, 0.0)\n        print(\"Epoch 0: BERT frozen\")\n    elif epoch == 1:\n        set_bert_lr(optimizer, 2e-5)\n        print(\"Epoch 1+: BERT unfrozen\")\n\n    model.train()\n    total_loss = 0\n\n    for batch in tqdm(loader, desc=f\"Epoch {epoch+1}/{NUM_EPOCHS}\"):\n\n        input_ids = batch[\"input_ids\"].to(device)\n        mask      = batch[\"attention_mask\"].to(device)\n        tokens    = batch[\"tokens\"].to(device)   # (B, 6, MAX_LEN)\n\n        optimizer.zero_grad()\n\n        batch_loss = 0\n\n        for layer_idx in range(NUM_LAYERS):\n            tgt    = tokens[:, layer_idx, :]\n            inp    = tgt[:, :-1]\n            target = tgt[:, 1:]\n\n            logits = model(input_ids, mask, inp, layer_idx)  # (B, S, V)\n\n            loss = criterion(\n                logits.reshape(-1, VOCAB_SIZE),\n                target.reshape(-1)\n            )\n            batch_loss += loss\n\n        batch_loss = batch_loss / NUM_LAYERS\n        batch_loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n\n        optimizer.step()\n        scheduler.step()\n\n        total_loss += batch_loss.item()\n\n    print(f\"Epoch {epoch+1} loss: {total_loss:.4f}\")\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T13:59:11.220124Z","iopub.execute_input":"2026-02-26T13:59:11.220812Z","iopub.status.idle":"2026-02-26T16:45:52.582761Z","shell.execute_reply.started":"2026-02-26T13:59:11.220781Z","shell.execute_reply":"2026-02-26T16:45:52.582067Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Smarter Inference with Gloss-Based Length + Temperature Sampling","metadata":{}},{"cell_type":"code","source":"# Build a lookup: gloss word count → median motion length\ntrain_df[\"gloss_words\"] = train_df[\"gloss\"].apply(lambda x: len(str(x).split()))\nlength_lookup = train_df.groupby(\"gloss_words\")[\"length\"].median().to_dict()\nglobal_median = int(train_df[\"length\"].median())\n\ndef predict_length(gloss_text):\n    \"\"\"Estimate motion length from gloss word count.\"\"\"\n    n_words = len(str(gloss_text).split())\n    length = length_lookup.get(n_words, global_median)\n    return int(np.clip(length, 40, 800))\n\ndef generate_tokens_fast(model, text, target_len, temperature=1.0, top_k=50):\n    \"\"\"\n    Generates all 6 RVQ layers.\n    - temperature > 1 → more diverse (helps Diversity metric)\n    - top_k sampling for quality + diversity balance\n    \"\"\"\n    model.eval()\n    \n    enc = tokenizer(\n        text,\n        padding=\"max_length\",\n        truncation=True,\n        max_length=64,\n        return_tensors=\"pt\"\n    ).to(device)\n    \n    all_layers = []\n    \n    with torch.no_grad():\n        memory = model.encode_text(enc[\"input_ids\"], enc[\"attention_mask\"])\n        \n        for layer_idx in range(NUM_LAYERS):\n            \n            generated = torch.zeros(1, 1, dtype=torch.long, device=device)\n            \n            for step in range(target_len):\n                inp = generated\n                \n                tgt_emb = model.token_embed(inp)\n                layer_id = torch.tensor([layer_idx], device=device)\n                layer_vec = model.layer_embed(layer_id).unsqueeze(1)\n                tgt_emb = tgt_emb + layer_vec\n                \n                S = inp.size(1)\n                tgt_mask = torch.triu(torch.ones(S, S, device=device), diagonal=1).bool()\n                \n                dec_out = model.decoder(tgt=tgt_emb, memory=memory, tgt_mask=tgt_mask)\n                logits = model.output_heads[layer_idx](dec_out[:, -1])  # (1, VOCAB_SIZE)\n                \n                # Top-k sampling\n                if temperature != 1.0:\n                    logits = logits / temperature\n                \n                if top_k > 0:\n                    top_vals, top_idx = torch.topk(logits, top_k, dim=-1)\n                    probs = torch.softmax(top_vals, dim=-1)\n                    sampled = torch.multinomial(probs, 1)  # (1, 1)\n                    next_token = top_idx.gather(-1, sampled)\n                else:\n                    next_token = torch.argmax(logits, dim=-1, keepdim=True)\n                \n                generated = torch.cat([generated, next_token], dim=1)\n            \n            all_layers.append(generated[:, 1:].squeeze(0).cpu().numpy())\n    \n    return all_layers\n\n\n# ---- Generate Submission ----\nmodel.eval()\n\nsubmission = []\nTEMPERATURE = 1.1   # Slightly above 1 boosts diversity score\n\nfor _, row in tqdm(test_df.iterrows(), total=len(test_df)):\n    \n    text = str(row[\"gloss\"]) + \" [SEP] \" + str(row[\"sentence\"])\n    target_len = predict_length(row[\"gloss\"])\n    \n    tokens = generate_tokens_fast(model, text, target_len, temperature=TEMPERATURE, top_k=40)\n    \n    row_dict = {\"id\": row[\"id\"]}\n    names = ['base_tokens','residual_1','residual_2','residual_3','residual_4','residual_5']\n    for i, name in enumerate(names):\n        row_dict[name] = \" \".join(map(str, tokens[i].tolist()))\n    \n    submission.append(row_dict)\n\nsubmission_df = pd.DataFrame(submission)\nsubmission_df.to_csv(\"submission.csv\", index=False)\nprint(\"Done!\", submission_df.shape)\nsubmission_df.head(2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T16:45:52.584339Z","iopub.execute_input":"2026-02-26T16:45:52.584864Z","iopub.status.idle":"2026-02-26T18:03:53.903644Z","shell.execute_reply.started":"2026-02-26T16:45:52.584834Z","shell.execute_reply":"2026-02-26T18:03:53.902300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}