{"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":[{"sourceId":130287,"databundleVersionId":15633993,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":82982,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":69710,"modelId":94840}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\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, Subset\nfrom transformers import T5Tokenizer, T5EncoderModel\nfrom tqdm.auto import tqdm\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport gc\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:50.388358Z","iopub.execute_input":"2026-02-17T14:22:50.388735Z","iopub.status.idle":"2026-02-17T14:22:54.698949Z","shell.execute_reply.started":"2026-02-17T14:22:50.388707Z","shell.execute_reply":"2026-02-17T14:22:54.697302Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\ntorch.manual_seed(SEED)\nnp.random.seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.699687Z","iopub.status.idle":"2026-02-17T14:22:54.699956Z","shell.execute_reply.started":"2026-02-17T14:22:54.699825Z","shell.execute_reply":"2026-02-17T14:22:54.699840Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CONFIG:\n    DATA_DIR = \"/kaggle/input/motion-s-hierarchical-text-to-motion-generation-for-sign-language\"\n    T5_PATH = \"/kaggle/input/models/hritik619916/t5-small/tensorflow2/default/1\"\n    OUTPUT_DIR = \"/kaggle/working\"\n    \n    TOKEN_VOCAB_SIZE = 512\n    MIN_LEN = 40\n    MAX_LEN = 800\n    MAX_TEXT_LEN = 128\n    \n    HIDDEN_DIM = 512\n    NUM_LAYERS = 6\n    NUM_HEADS = 8\n    DROPOUT = 0.1\n    \n    BATCH_SIZE = 32\n    ACCUM_STEPS = 2\n    LR = 3e-4\n    WARMUP_STEPS = 500\n    TOTAL_EPOCHS = 10\n    WEIGHT_DECAY = 0.01\n    EARLY_STOPPING_PATIENCE = 3\n    \n    BASE_WEIGHT = 1.0\n    RESIDUAL_WEIGHTS = [0.8, 0.6, 0.4, 0.3, 0.2]\n    \n    TEACHER_FORCING_START = 0.8\n    TEACHER_FORCING_END = 0.2\n    \n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    USE_FP16 = True\n    \n    SHIFT_TOKENS = True   # will be updated after EDA\n\nprint(f\" Device: {CONFIG.DEVICE} | FP16: {CONFIG.USE_FP16}\")\nprint(f\" Batch size: {CONFIG.BATCH_SIZE} (effective: {CONFIG.BATCH_SIZE * CONFIG.ACCUM_STEPS})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.700960Z","iopub.status.idle":"2026-02-17T14:22:54.701330Z","shell.execute_reply.started":"2026-02-17T14:22:54.701141Z","shell.execute_reply":"2026-02-17T14:22:54.701162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(f\"{CONFIG.DATA_DIR}/train.csv\")\nprint(f\"Total samples: {len(df)}\")\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.702661Z","iopub.status.idle":"2026-02-17T14:22:54.703049Z","shell.execute_reply.started":"2026-02-17T14:22:54.702850Z","shell.execute_reply":"2026-02-17T14:22:54.702876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check if token 0 exists (to decide on padding ID)\ntoken_columns = ['base_tokens', 'residual_1', 'residual_2', 'residual_3', 'residual_4', 'residual_5']\nall_tokens = []\nfor col in token_columns:\n    tokens_series = df[col].dropna().str.split()\n    for tokens in tokens_series:\n        all_tokens.extend([int(t) for t in tokens])\nunique_tokens = set(all_tokens)\nprint(f\"Unique token IDs: {sorted(unique_tokens)[:20]} ... (total {len(unique_tokens)})\")\nprint(f\"Is token 0 present? {0 in unique_tokens}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.704242Z","iopub.status.idle":"2026-02-17T14:22:54.704647Z","shell.execute_reply.started":"2026-02-17T14:22:54.704438Z","shell.execute_reply":"2026-02-17T14:22:54.704461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# If token 0 present, we must shift all tokens by +1 and use 0 for padding\nCONFIG.SHIFT_TOKENS = 0 in unique_tokens\nif CONFIG.SHIFT_TOKENS:\n    CONFIG.TOKEN_VOCAB_SIZE = 513\n    print(f\"Adjusted vocab size to {CONFIG.TOKEN_VOCAB_SIZE} (including padding ID 0)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.705718Z","iopub.status.idle":"2026-02-17T14:22:54.705970Z","shell.execute_reply.started":"2026-02-17T14:22:54.705853Z","shell.execute_reply":"2026-02-17T14:22:54.705868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SignLanguageDataset(Dataset):\n    def __init__(self, csv_path, split='train', shift_tokens=True, pre_tokenize=True):\n        self.df = pd.read_csv(csv_path)\n        self.shift_tokens = shift_tokens\n        self.texts = self.df['gloss'].fillna(self.df['sentence']).astype(str).values\n        \n        if pre_tokenize:\n            print(\"⚡ Pre-tokenizing dataset...\")\n            self.tokens = []\n            for layer in token_columns:\n                token_lists = self.df[layer].str.split().apply(\n                    lambda x: np.array(x, dtype=np.int16) if isinstance(x, list) else np.array([], dtype=np.int16)\n                ).values\n                self.tokens.append(token_lists)\n            \n            if shift_tokens:\n                for i in range(len(self.tokens)):\n                    for j in range(len(self.tokens[i])):\n                        if len(self.tokens[i][j]) > 0:\n                            self.tokens[i][j] = self.tokens[i][j] + 1\n            \n            # Identify valid samples (length within bounds and consistent across layers)\n            base_lengths = np.array([len(tok) for tok in self.tokens[0]])\n            valid_mask = (base_lengths >= CONFIG.MIN_LEN) & (base_lengths <= CONFIG.MAX_LEN)\n            \n            # Check layer length consistency\n            for i in range(1, len(self.tokens)):\n                layer_lengths = np.array([len(tok) for tok in self.tokens[i]])\n                valid_mask &= (layer_lengths == base_lengths)\n            \n            self.valid_indices = np.where(valid_mask)[0]\n            print(f\"✅ Pre-tokenization complete: {len(self.df)} total, {len(self.valid_indices)} valid samples\")\n    \n    def __len__(self):\n        return len(self.valid_indices) if hasattr(self, 'valid_indices') else len(self.df)\n    \n    def __getitem__(self, idx):\n        # Map idx to original index\n        real_idx = self.valid_indices[idx] if hasattr(self, 'valid_indices') else idx\n        text = self.texts[real_idx]\n        length = len(self.tokens[0][real_idx])\n        \n        # No assertion needed now; already filtered\n        return {\n            'text': text,\n            'length': length,\n            'base_tokens': torch.tensor(self.tokens[0][real_idx], dtype=torch.long),\n            'residuals': [torch.tensor(self.tokens[i+1][real_idx], dtype=torch.long) for i in range(5)]\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.707537Z","iopub.status.idle":"2026-02-17T14:22:54.708254Z","shell.execute_reply.started":"2026-02-17T14:22:54.708043Z","shell.execute_reply":"2026-02-17T14:22:54.708068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def sign_collate_fn(batch):\n    texts = [item['text'] for item in batch]\n    lengths = torch.tensor([item['length'] for item in batch], dtype=torch.long)\n    max_len = lengths.max().item()\n    \n    base_tokens = torch.zeros(len(batch), max_len, dtype=torch.long)\n    residuals = [torch.zeros(len(batch), max_len, dtype=torch.long) for _ in range(5)]\n    for i, item in enumerate(batch):\n        l = item['length']\n        base_tokens[i, :l] = item['base_tokens']\n        for j in range(5):\n            residuals[j][i, :l] = item['residuals'][j]\n    \n    padding_mask = torch.zeros(len(batch), max_len, dtype=torch.bool)\n    for i, l in enumerate(lengths):\n        padding_mask[i, l:] = True\n    \n    return {\n        'text': texts,\n        'length': lengths,\n        'base_tokens': base_tokens,\n        'residuals': residuals,\n        'padding_mask': padding_mask\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.709675Z","iopub.status.idle":"2026-02-17T14:22:54.710066Z","shell.execute_reply.started":"2026-02-17T14:22:54.709883Z","shell.execute_reply":"2026-02-17T14:22:54.709905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=1000):\n        super().__init__()\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() * (-np.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe.unsqueeze(0))\n    \n    def forward(self, x, seq_len=None):\n        if seq_len is None:\n            seq_len = x.size(1)\n        return x + self.pe[:, :seq_len, :]\n\nclass BaseTokenGenerator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.text_proj = nn.Linear(512, CONFIG.HIDDEN_DIM)\n        self.pos_enc = PositionalEncoding(CONFIG.HIDDEN_DIM, max_len=CONFIG.MAX_LEN)\n        decoder_layer = nn.TransformerDecoderLayer(\n            d_model=CONFIG.HIDDEN_DIM,\n            nhead=CONFIG.NUM_HEADS,\n            dim_feedforward=2048,\n            dropout=CONFIG.DROPOUT,\n            batch_first=True,\n            norm_first=True\n        )\n        self.transformer = nn.TransformerDecoder(decoder_layer, num_layers=CONFIG.NUM_LAYERS)\n        self.out_proj = nn.Linear(CONFIG.HIDDEN_DIM, CONFIG.TOKEN_VOCAB_SIZE)\n        self.layer_norm = nn.LayerNorm(CONFIG.HIDDEN_DIM)\n        self.use_checkpointing = True\n    \n    def _generate_square_subsequent_mask(self, sz):\n        return torch.triu(torch.ones(sz, sz), diagonal=1).bool()\n    \n    def forward(self, text_emb, target_lengths):\n        B = text_emb.size(0)\n        L_max = target_lengths.max().item()\n        text_proj = self.text_proj(text_emb)\n        tgt = torch.zeros(B, L_max, CONFIG.HIDDEN_DIM, device=text_emb.device)\n        tgt = self.pos_enc(tgt, seq_len=L_max)\n        causal_mask = self._generate_square_subsequent_mask(L_max).to(text_emb.device)\n        padding_mask = torch.zeros(B, L_max, device=text_emb.device, dtype=torch.bool)\n        for i, l in enumerate(target_lengths):\n            padding_mask[i, l:] = True\n        \n        if self.use_checkpointing and self.training:\n            tgt = torch.utils.checkpoint.checkpoint(\n                self.transformer, tgt, text_proj, causal_mask, None, padding_mask, use_reentrant=False\n            )\n        else:\n            tgt = self.transformer(tgt, text_proj, tgt_mask=causal_mask, tgt_key_padding_mask=padding_mask)\n        tgt = self.layer_norm(tgt)\n        logits = self.out_proj(tgt)\n        return logits, padding_mask\n\nclass ResidualTokenGenerator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.base_emb = nn.Embedding(CONFIG.TOKEN_VOCAB_SIZE, CONFIG.HIDDEN_DIM // 2, padding_idx=0)\n        self.text_proj = nn.Linear(512, CONFIG.HIDDEN_DIM // 2)\n        self.pos_enc = PositionalEncoding(CONFIG.HIDDEN_DIM, max_len=CONFIG.MAX_LEN)\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=CONFIG.HIDDEN_DIM,\n            nhead=CONFIG.NUM_HEADS,\n            dim_feedforward=2048,\n            dropout=CONFIG.DROPOUT,\n            batch_first=True,\n            norm_first=True\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=CONFIG.NUM_LAYERS)\n        self.heads = nn.ModuleList([\n            nn.Sequential(\n                nn.Linear(CONFIG.HIDDEN_DIM, CONFIG.HIDDEN_DIM),\n                nn.GELU(),\n                nn.LayerNorm(CONFIG.HIDDEN_DIM),\n                nn.Linear(CONFIG.HIDDEN_DIM, CONFIG.TOKEN_VOCAB_SIZE)\n            ) for _ in range(5)\n        ])\n        # Initialize deeper residuals with smaller weights\n        for i, head in enumerate(self.heads):\n            if i > 0:\n                for param in head[-1].parameters():\n                    param.data *= 0.5\n    \n    def forward(self, base_tokens, text_emb, padding_mask, use_teacher_forcing=False, ground_truth=None):\n        if use_teacher_forcing and ground_truth is not None:\n            base_input = ground_truth\n        else:\n            base_input = base_tokens  # must be provided when not teacher forcing\n        \n        B, L = base_input.size()\n        base_emb = self.base_emb(base_input)\n        with torch.no_grad():\n            text_pooled = text_emb.mean(dim=1)\n        text_proj = self.text_proj(text_pooled).unsqueeze(1).expand(-1, L, -1)\n        tgt = torch.cat([base_emb, text_proj], dim=-1)\n        tgt = self.pos_enc(tgt, seq_len=L)\n        tgt = self.transformer(tgt, src_key_padding_mask=padding_mask)\n        residuals = [head(tgt) for head in self.heads]\n        return residuals","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.711140Z","iopub.status.idle":"2026-02-17T14:22:54.711451Z","shell.execute_reply.started":"2026-02-17T14:22:54.711283Z","shell.execute_reply":"2026-02-17T14:22:54.711299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TrainingEngine:\n    def __init__(self):\n        self.tokenizer = T5Tokenizer.from_pretrained(CONFIG.T5_PATH, local_files_only=True, use_fast=False)\n        self.text_encoder = T5EncoderModel.from_pretrained(CONFIG.T5_PATH, local_files_only=True).to(CONFIG.DEVICE).eval()\n        \n        self.base_gen = BaseTokenGenerator().to(CONFIG.DEVICE)\n        self.res_gen = ResidualTokenGenerator().to(CONFIG.DEVICE)\n        \n        self.optimizer = torch.optim.AdamW([\n            {'params': self.base_gen.parameters(), 'lr': CONFIG.LR * 0.8},\n            {'params': self.res_gen.parameters(), 'lr': CONFIG.LR}\n        ], weight_decay=CONFIG.WEIGHT_DECAY)\n        \n        self.scheduler = torch.optim.lr_scheduler.LambdaLR(\n            self.optimizer,\n            lr_lambda=lambda step: min(1.0, step / CONFIG.WARMUP_STEPS) if step < CONFIG.WARMUP_STEPS else\n            0.5 * (1 + np.cos(np.pi * (step - CONFIG.WARMUP_STEPS) / (CONFIG.TOTAL_EPOCHS * 1000)))\n        )\n        \n        self.scaler = torch.cuda.amp.GradScaler(enabled=CONFIG.USE_FP16)\n        self.criterion = nn.CrossEntropyLoss(label_smoothing=0.1, ignore_index=0)\n    \n    def encode_text(self, texts):\n        with torch.no_grad():\n            inputs = self.tokenizer(texts, return_tensors=\"pt\", padding=True, truncation=True,\n                                    max_length=CONFIG.MAX_TEXT_LEN).to(CONFIG.DEVICE)\n            outputs = self.text_encoder(**inputs)\n        return outputs.last_hidden_state\n    \n    def compute_loss_and_acc(self, logits, targets, padding_mask):\n        B, L, V = logits.size()\n        logits_flat = logits.view(-1, V)\n        targets_flat = targets.view(-1)\n        mask_flat = ~padding_mask.view(-1)\n        loss = self.criterion(logits_flat, targets_flat)\n        preds = logits_flat.argmax(dim=-1)\n        correct = (preds == targets_flat) & mask_flat\n        acc = correct.sum().float() / mask_flat.sum().float() if mask_flat.sum() > 0 else torch.tensor(0.0)\n        return loss, acc.item()\n    \n    def train_epoch(self, dataloader, epoch):\n        self.base_gen.train()\n        self.res_gen.train()\n        total_loss = total_base_loss = total_base_acc = 0\n        total_res_loss = [0]*5\n        total_res_acc = [0]*5\n        teacher_forcing_ratio = CONFIG.TEACHER_FORCING_START - (\n            (CONFIG.TEACHER_FORCING_START - CONFIG.TEACHER_FORCING_END) * epoch / (CONFIG.TOTAL_EPOCHS - 1)\n        )\n        \n        progress = tqdm(dataloader, desc=f\"Epoch {epoch+1}/{CONFIG.TOTAL_EPOCHS}\")\n        for step, batch in enumerate(progress):\n            texts = batch['text']\n            lengths = batch['length'].to(CONFIG.DEVICE)\n            base_targets = batch['base_tokens'].to(CONFIG.DEVICE)\n            res_targets = [t.to(CONFIG.DEVICE) for t in batch['residuals']]\n            padding_mask = batch['padding_mask'].to(CONFIG.DEVICE)\n            \n            text_emb = self.encode_text(texts)\n            \n            with torch.cuda.amp.autocast(enabled=CONFIG.USE_FP16):\n                base_logits, base_padding_mask = self.base_gen(text_emb, lengths)\n                use_tf = np.random.random() < teacher_forcing_ratio\n                if use_tf:\n                    res_logits = self.res_gen(None, text_emb, base_padding_mask,\n                                              use_teacher_forcing=True, ground_truth=base_targets)\n                else:\n                    with torch.no_grad():\n                        probs = F.softmax(base_logits / 0.8, dim=-1)\n                        base_preds = torch.multinomial(probs.view(-1, probs.size(-1)), 1).view(base_logits.size(0), base_logits.size(1))\n                    res_logits = self.res_gen(base_preds, text_emb, base_padding_mask, use_teacher_forcing=False)\n                \n                base_loss, base_acc = self.compute_loss_and_acc(base_logits, base_targets, base_padding_mask)\n                res_losses, res_accs = [], []\n                for i in range(5):\n                    l, a = self.compute_loss_and_acc(res_logits[i], res_targets[i], base_padding_mask)\n                    res_losses.append(l); res_accs.append(a)\n                \n                total = CONFIG.BASE_WEIGHT * base_loss + sum(w*l for w,l in zip(CONFIG.RESIDUAL_WEIGHTS, res_losses))\n            \n            self.scaler.scale(total).backward()\n            if (step+1) % CONFIG.ACCUM_STEPS == 0:\n                self.scaler.unscale_(self.optimizer)\n                torch.nn.utils.clip_grad_norm_(list(self.base_gen.parameters())+list(self.res_gen.parameters()), max_norm=1.0)\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n                self.optimizer.zero_grad()\n                self.scheduler.step()\n            \n            total_loss += total.item()\n            total_base_loss += base_loss.item()\n            total_base_acc += base_acc\n            for i in range(5):\n                total_res_loss[i] += res_losses[i].item()\n                total_res_acc[i] += res_accs[i]\n            \n            progress.set_postfix({'loss': f\"{total_loss/(step+1):.4f}\", 'base_acc': f\"{total_base_acc/(step+1):.4f}\"})\n        \n        n = len(dataloader)\n        return {'total_loss': total_loss/n, 'base_loss': total_base_loss/n, 'base_acc': total_base_acc/n,\n                'res_losses': [l/n for l in total_res_loss], 'res_accs': [a/n for a in total_res_acc],\n                'teacher_forcing_ratio': teacher_forcing_ratio}\n    \n    def validate(self, dataloader):\n        self.base_gen.eval(); self.res_gen.eval()\n        total_loss = total_base_acc = 0\n        with torch.no_grad():\n            for batch in tqdm(dataloader, desc=\"Validation\", leave=False):\n                texts = batch['text']\n                lengths = batch['length'].to(CONFIG.DEVICE)\n                base_targets = batch['base_tokens'].to(CONFIG.DEVICE)\n                res_targets = [t.to(CONFIG.DEVICE) for t in batch['residuals']]\n                padding_mask = batch['padding_mask'].to(CONFIG.DEVICE)\n                text_emb = self.encode_text(texts)\n                base_logits, base_padding_mask = self.base_gen(text_emb, lengths)\n                base_preds = base_logits.argmax(dim=-1)\n                res_logits = self.res_gen(base_preds, text_emb, base_padding_mask, use_teacher_forcing=False)\n                base_loss, base_acc = self.compute_loss_and_acc(base_logits, base_targets, base_padding_mask)\n                res_losses = [self.compute_loss_and_acc(res_logits[i], res_targets[i], base_padding_mask)[0] for i in range(5)]\n                total = base_loss + sum(res_losses)/5\n                total_loss += total.item()\n                total_base_acc += base_acc\n        n = len(dataloader)\n        return total_loss/n, total_base_acc/n\n    \n    def save_checkpoint(self, epoch, metrics, path):\n        torch.save({'epoch': epoch, 'base_gen_state': self.base_gen.state_dict(),\n                    'res_gen_state': self.res_gen.state_dict(), 'optimizer_state': self.optimizer.state_dict(),\n                    'scheduler_state': self.scheduler.state_dict(), 'metrics': metrics}, path)\n        print(f\" Checkpoint saved: {path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.712514Z","iopub.status.idle":"2026-02-17T14:22:54.712841Z","shell.execute_reply.started":"2026-02-17T14:22:54.712721Z","shell.execute_reply":"2026-02-17T14:22:54.712737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nfull_dataset = SignLanguageDataset(f\"{CONFIG.DATA_DIR}/train.csv\", shift_tokens=CONFIG.SHIFT_TOKENS, pre_tokenize=True)\nn_train = int(len(full_dataset) * 0.95)\ntrain_ds = Subset(full_dataset, list(range(n_train)))\nval_ds = Subset(full_dataset, list(range(n_train, len(full_dataset))))\nprint(f\"Training samples: {len(train_ds):,} | Validation samples: {len(val_ds):,}\")\n\ntrain_loader = DataLoader(train_ds, batch_size=CONFIG.BATCH_SIZE, shuffle=True, num_workers=2,\n                          pin_memory=True, collate_fn=sign_collate_fn, drop_last=True)\nval_loader = DataLoader(val_ds, batch_size=CONFIG.BATCH_SIZE*2, shuffle=False, num_workers=2,\n                        pin_memory=True, collate_fn=sign_collate_fn)\n\nengine = TrainingEngine()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.715920Z","iopub.status.idle":"2026-02-17T14:22:54.716179Z","shell.execute_reply.started":"2026-02-17T14:22:54.716063Z","shell.execute_reply":"2026-02-17T14:22:54.716078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_val_loss = float('inf')\npatience_counter = 0\n\nfor epoch in range(CONFIG.TOTAL_EPOCHS):\n    print(f\"\\n{'='*70}\\nEpoch {epoch+1}/{CONFIG.TOTAL_EPOCHS}\\n{'='*70}\")\n    train_metrics = engine.train_epoch(train_loader, epoch)\n    print(\"\\n Running validation...\")\n    val_loss, val_base_acc = engine.validate(val_loader)\n    \n    print(f\"\\n Epoch {epoch+1} Complete\")\n    print(f\"   Train Loss: {train_metrics['total_loss']:.4f} | Val Loss: {val_loss:.4f}\")\n    print(f\"   Train Base Acc: {train_metrics['base_acc']:.4f} | Val Base Acc: {val_base_acc:.4f}\")\n    \n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        patience_counter = 0\n        checkpoint_path = f\"{CONFIG.OUTPUT_DIR}/best_model_epoch{epoch+1}.pth\"\n        engine.save_checkpoint(epoch, {'val_loss': val_loss, 'val_acc': val_base_acc}, checkpoint_path)\n        print(f\"    New best model saved\")\n    else:\n        patience_counter += 1\n        print(f\"    No improvement. Patience: {patience_counter}/{CONFIG.EARLY_STOPPING_PATIENCE}\")\n        if patience_counter >= CONFIG.EARLY_STOPPING_PATIENCE:\n            print(\"    Early stopping triggered.\")\n            break\n\nengine.save_checkpoint(epoch, {'val_loss': val_loss}, f\"{CONFIG.OUTPUT_DIR}/final_model.pth\")\nprint(f\"\\n Final model saved.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.717387Z","iopub.status.idle":"2026-02-17T14:22:54.717739Z","shell.execute_reply.started":"2026-02-17T14:22:54.717601Z","shell.execute_reply":"2026-02-17T14:22:54.717624Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_csv_path = f\"{CONFIG.DATA_DIR}/test.csv\"\nif os.path.exists(test_csv_path):\n    \n    # Find the best checkpoint (latest epoch with lowest val loss)\n    best_checkpoint = None\n    best_epoch = None\n    for f in os.listdir(CONFIG.OUTPUT_DIR):\n        if f.startswith(\"best_model_epoch\") and f.endswith(\".pth\"):\n            epoch_num = int(f.split('_')[-1].split('.')[0].replace('epoch',''))\n            if best_epoch is None or epoch_num > best_epoch:\n                best_epoch = epoch_num\n                best_checkpoint = os.path.join(CONFIG.OUTPUT_DIR, f)\n    if best_checkpoint is None:\n        best_checkpoint = f\"{CONFIG.OUTPUT_DIR}/final_model.pth\"\n    print(f\"Using checkpoint: {best_checkpoint}\")\n    \n    # Load model\n    checkpoint = torch.load(best_checkpoint, map_location=CONFIG.DEVICE, weights_only=False)\n    base_gen = BaseTokenGenerator().to(CONFIG.DEVICE)\n    res_gen = ResidualTokenGenerator().to(CONFIG.DEVICE)\n    base_gen.load_state_dict(checkpoint['base_gen_state'])\n    res_gen.load_state_dict(checkpoint['res_gen_state'])\n    base_gen.eval()\n    res_gen.eval()\n    \n\n    # Load test data\n    test_df = pd.read_csv(test_csv_path)\n    test_texts = test_df['gloss'].fillna(test_df['sentence']).astype(str).values\n    \n    # Check if 'id' column exists, otherwise use index\n    if 'id' in test_df.columns:\n        sample_ids = test_df['id'].values\n        print(\"Using 'id' column from test.csv\")\n    else:\n        sample_ids = test_df.index.values\n        print(\"No 'id' column found, using index as sample_id\")\n\n    # We need length for each test sample. Use a simple heuristic (median length from training)\n    # In production, you might use the official length predictor. Here we use median.\n    train_lengths = [len(full_dataset.tokens[0][i]) for i in range(len(full_dataset))]\n    median_length = int(np.median(train_lengths))\n    print(f\"Using median length from training: {median_length}\")\n    \n    # Inference dataset (without labels)\n    class TestDataset(Dataset):\n        def __init__(self, texts, length):\n            self.texts = texts\n            self.length = length\n        def __len__(self):\n            return len(self.texts)\n        def __getitem__(self, idx):\n            return {'text': self.texts[idx], 'length': self.length}\n    \n    def test_collate_fn(batch):\n        texts = [item['text'] for item in batch]\n        lengths = torch.tensor([item['length'] for item in batch], dtype=torch.long)\n        max_len = lengths.max().item()\n        padding_mask = torch.zeros(len(batch), max_len, dtype=torch.bool)\n        for i, l in enumerate(lengths):\n            padding_mask[i, l:] = True\n        return {'text': texts, 'length': lengths, 'padding_mask': padding_mask}\n    \n    test_ds = TestDataset(test_texts, median_length)\n    test_loader = DataLoader(test_ds, batch_size=CONFIG.BATCH_SIZE*2, shuffle=False,\n                             num_workers=2, pin_memory=True, collate_fn=test_collate_fn)\n    \n    # Generate predictions\n    all_base = []\n    all_res = [[] for _ in range(5)]\n    \n    tokenizer = T5Tokenizer.from_pretrained(CONFIG.T5_PATH, local_files_only=True, use_fast=False)\n    text_encoder = T5EncoderModel.from_pretrained(CONFIG.T5_PATH, local_files_only=True).to(CONFIG.DEVICE).eval()\n    \n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc=\"Inference\"):\n            texts = batch['text']\n            lengths = batch['length'].to(CONFIG.DEVICE)\n            padding_mask = batch['padding_mask'].to(CONFIG.DEVICE)\n            \n            inputs = tokenizer(texts, return_tensors=\"pt\", padding=True, truncation=True,\n                               max_length=CONFIG.MAX_TEXT_LEN).to(CONFIG.DEVICE)\n            text_emb = text_encoder(**inputs).last_hidden_state\n            \n            base_logits, _ = base_gen(text_emb, lengths)\n            base_preds = base_logits.argmax(dim=-1)  # [B, L]\n            \n            res_logits = res_gen(base_preds, text_emb, padding_mask, use_teacher_forcing=False)\n            res_preds = [r.argmax(dim=-1) for r in res_logits]\n            \n            all_base.append(base_preds.cpu())\n            for i in range(5):\n                all_res[i].append(res_preds[i].cpu())\n    \n    # Concatenate batches\n    base_tokens = torch.cat(all_base, dim=0).numpy()\n    res_tokens = [torch.cat(all_res[i], dim=0).numpy() for i in range(5)]\n    \n    # Revert token shift (subtract 1) and clip to [0,511]\n    if CONFIG.SHIFT_TOKENS:\n        base_tokens = base_tokens - 1\n        res_tokens = [r - 1 for r in res_tokens]\n    base_tokens = np.clip(base_tokens, 0, 511)\n    res_tokens = [np.clip(r, 0, 511) for r in res_tokens]\n    \n    # Convert each sample's tokens to space-separated strings\n    # Convert each sample's tokens to space-separated strings\n    submission = pd.DataFrame()\n    submission['id'] = sample_ids  # Use the correct column name 'id'\n    submission['base_tokens'] = [' '.join(map(str, row)) for row in base_tokens]\n    for i in range(5):\n        submission[f'residual_{i+1}'] = [' '.join(map(str, row)) for row in res_tokens[i]]\n    \n    # Sort by id to ensure correct order (if needed)\n    submission = submission.sort_values('id').reset_index(drop=True)\n    \n    # Save submission\n    submission_path = f\"{CONFIG.OUTPUT_DIR}/submission.csv\"\n    submission.to_csv(submission_path, index=False)\n    print(f\"\\n✅ Submission saved: {submission_path}\")\n    print(submission.head())\nelse:\n    print(\"\\n test.csv not found. Skipping submission generation.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.719187Z","iopub.status.idle":"2026-02-17T14:22:54.719419Z","shell.execute_reply.started":"2026-02-17T14:22:54.719309Z","shell.execute_reply":"2026-02-17T14:22:54.719322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T14:22:54.720449Z","iopub.status.idle":"2026-02-17T14:22:54.720776Z","shell.execute_reply.started":"2026-02-17T14:22:54.720612Z","shell.execute_reply":"2026-02-17T14:22:54.720635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}