{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":130287,"databundleVersionId":15633993},{"sourceType":"modelInstanceVersion","sourceId":772994,"databundleVersionId":15917634,"modelInstanceId":590360,"modelId":602693}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Task\nGenerate the 6 token layers given glossified text input.\n\n## How it works:\n\nEncoding: Motion sequence → 6 token layers (hierarchical from coarse to fine)\n\nDecoding: 6 token layers → Reconstructed motion sequence\n\n## Recommended Approach:\n\n- Start with the baseline notebook\n- Experiment with different text encoders (BERT, CLIP, T5, etc.)\n- Try autoregressive vs. non-autoregressive generation\n- Fine-tune on the provided training data\n- Validate locally before submitting\n\n## Requirements:\n\n- 6 Columns: You must provide all 6 RVQ layers.\n- Space-Separated: Each cell must contain token indices separated by a single space (do not use commas).\n- Token Values: Every token must be an integer .\n- Sequence Length: Each sequence must be between 40 and 800 tokens long.\n- Layer Consistency: All 6 layers must have the exact same sequence length for any given row.\n\n## Pipeline\n- Input: Glossified text (Example: \"ME GO STORE TOMORROW\")\n- Output: 6 sequences of integers `[0,511]` | All same length | Length between `[40,800]`\n\n## Model we're building\n\n0. Length Estimator\n1. Text Encoder -> Convert text into numeric representation (Text -> embedding vector)\n2. Token Generator -> Generate 6 layers where Each layer is of length K and Each token ∈ `[0, 511]`\n\n## Flowchart\n\n```mermaid\nflowchart TD\n    A[Gloss Text Input] --> B[Text Encoder]\n    B --> C[Length Estimator]\n    B --> D[Token Generator Transformer]\n    C --> D\n    D --> E[Base Tokens 0 to 511]\n    D --> F[Residual Layer 1]\n    D --> G[Residual Layer 2]\n    D --> H[Residual Layer 3]\n    D --> I[Residual Layer 4]\n    D --> J[Residual Layer 5]\n    E --> K[Submission CSV]\n    F --> K\n    G --> K\n    H --> K\n    I --> K\n    J --> K\n```","metadata":{"_cell_guid":"c6127970-3f05-4311-ae82-6f2c40db0c17","_uuid":"e1032a01-6e22-4d65-b7be-ad9b8c7d4be4","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"!date\nimport os\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\nimport sys\nfrom pathlib import Path\nfrom collections import Counter, defaultdict\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import T5Tokenizer, T5EncoderModel\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm import tqdm\nimport warnings\nimport random\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nIS_KAGGLE = Path(\"/kaggle/working\").exists()\nINPUT_DIR = Path(\"/kaggle/input/motion-s-hierarchical-text-to-motion-generation-for-sign-language\")\nMODEL_DIR = Path(\"/kaggle/input/motion-s-vae-rvq/pytorch/default/3\")\nDATASET_ROOT = INPUT_DIR / \"Train\"\nMOTION_FEATS_DIR = INPUT_DIR / \"Motion-Features\"\nTRAIN_CSV = INPUT_DIR / \"train.csv\"\nTEST_CSV = INPUT_DIR / \"test.csv\"\nMODEL_PATH = MODEL_DIR / \"rvq_vae_best.pth\"\nNORM_PATH = INPUT_DIR / \"normalization.npz\"\nT5_PATH = \"/kaggle/input/models/pad1tya/t5-base/pytorch/default/1/t5-base-local/\"\nBEST_MODEL_PATH = \"/kaggle/working/best_model.pth\"\nLAST_MODEL_PATH = \"/kaggle/working/last_checkpoint.pth\"\nNOTEBOOK_EXPERIMENT_MODE = False\n\n\ndef get_bvh_path(inp):\n    \"\"\"Handles both raw IDs and the //dataset/ paths from the CSV.\"\"\"\n    inp_str = str(inp)\n    if \"dataset\" in inp_str:\n        clean_path = inp_str.replace(\"//dataset/\", \"\").replace(\"dataset/\", \"\").lstrip(\"/\")\n        return DATASET_ROOT / clean_path\n    return DATASET_ROOT / inp_str / f\"{inp_str}.bvh\"\n\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark     = False\n\nset_seed(42)\n\nprint(\"DATASET_ROOT:=\", DATASET_ROOT)\nprint(\"NOTEBOOK_EXPERIMENT_MODE\", NOTEBOOK_EXPERIMENT_MODE)\nassert DATASET_ROOT.exists()","metadata":{"_cell_guid":"2eec5a7c-f0f7-400a-8724-43423488c0cd","_uuid":"e5700200-a470-43dc-878c-19fe293ae4cd","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T10:40:20.656825Z","iopub.execute_input":"2026-03-04T10:40:20.657540Z","iopub.status.idle":"2026-03-04T10:40:20.868608Z","shell.execute_reply.started":"2026-03-04T10:40:20.657512Z","shell.execute_reply":"2026-03-04T10:40:20.867943Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Visualizations and Statistics","metadata":{}},{"cell_type":"code","source":"if NOTEBOOK_EXPERIMENT_MODE:\n    def load_metadata(sample_dir: Path) -> dict:\n        \"\"\"Load metadata.txt from a sample directory.\"\"\"\n        metadata_path = sample_dir / \"metadata.txt\"\n        if not metadata_path.exists():\n            return None\n        \n        result = {}\n        try:\n            with open(metadata_path, 'r', encoding='utf-8') as f:\n                for line in f:\n                    line = line.strip()\n                    if line.startswith(\"SENTENCE:\"):\n                        result['sentence'] = line.replace(\"SENTENCE:\", \"\").strip()\n                    elif line.startswith(\"GLOSS:\"):\n                        result['gloss'] = line.replace(\"GLOSS:\", \"\").strip()\n        except Exception as e:\n            return None\n        \n        return result if 'sentence' in result and 'gloss' in result else None\n    \n    \n    def count_bvh_files(sample_dir: Path) -> int:\n        \"\"\"Count BVH files in a sample directory.\"\"\"\n        return len(list(sample_dir.glob(\"*.bvh\")))\n    \n    \n    def parse_glosses(gloss_str: str) -> list:\n        \"\"\"Parse gloss string into list of individual glosses.\"\"\"\n        # Remove sentence boundary markers and split\n        cleaned = gloss_str.replace(\"//\", \"\").strip()\n        return [g.strip() for g in cleaned.split() if g.strip()]\n    \n    \n    def is_fingerspelling(gloss: str) -> bool:\n        \"\"\"Check if a gloss is a fingerspelled letter (single uppercase letter).\"\"\"\n        return len(gloss) == 1 and gloss.isupper()\n    \n    import os\n    import pandas as pd\n    from pathlib import Path\n    from tqdm import tqdm\n    \n    def sample_generator(dataset_root):\n        total = sum(1 for e in os.scandir(dataset_root) if e.is_dir())\n        with os.scandir(dataset_root) as it:\n            for entry in tqdm(it, total=total, desc=\"Loading samples\"):\n                if not entry.is_dir():\n                    continue\n    \n                sample_path = Path(entry.path)\n                sample_id = entry.name\n                metadata = load_metadata(sample_path)\n                bvh_count = count_bvh_files(sample_path)\n    \n                if metadata:\n                    sentence = metadata.get('sentence', '')\n                    gloss = metadata.get('gloss', '')\n                    glosses = parse_glosses(gloss)\n                    fingerspell_count = sum(is_fingerspelling(g) for g in glosses)\n    \n                    yield {\n                        'sample_id': sample_id,\n                        'sentence': sentence,\n                        'gloss': gloss,\n                        'gloss_list': glosses,\n                        'gloss_count': len(glosses),\n                        'fingerspell_count': fingerspell_count,\n                        'bvh_count': bvh_count,\n                        'sentence_word_count': len(sentence.split()),\n                        'sentence_char_count': len(sentence),\n                    }\n                else:\n                    yield {\n                        'sample_id': sample_id,\n                        'sentence': None,\n                        'gloss': None,\n                        'gloss_list': [],\n                        'gloss_count': 0,\n                        'fingerspell_count': 0,\n                        'bvh_count': bvh_count,\n                        'sentence_word_count': 0,\n                        'sentence_char_count': 0,\n                    }\n    \n    if DATASET_ROOT.exists():\n        df = pd.DataFrame(sample_generator(DATASET_ROOT))\n        print(f\"Loaded {len(df)} samples\")\n    \n    # Basic statistics\n    valid_df = df[df['sentence'].notna()]\n    \n    print(\"=\" * 50)\n    print(\"DATASET STATISTICS\")\n    print(\"=\" * 50)\n    print(f\"Total samples:           {len(df):,}\")\n    print(f\"Valid samples (w/ meta): {len(valid_df):,}\")\n    print(f\"Missing metadata:        {len(df) - len(valid_df):,}\")\n    print(f\"Total BVH files:         {df['bvh_count'].sum():,}\")\n    print()\n    print(\"BVH files per sample:\")\n    print(f\"  Mean:   {df['bvh_count'].mean():.1f}\")\n    print(f\"  Median: {df['bvh_count'].median():.1f}\")\n    print(f\"  Min:    {df['bvh_count'].min()}\")\n    print(f\"  Max:    {df['bvh_count'].max()}\")\n    print()\n    \n    # Sentence length statistics\n    print(\"SENTENCE STATISTICS\")\n    print(\"-\" * 40)\n    print(f\"Word count:\")\n    print(f\"  Mean:   {valid_df['sentence_word_count'].mean():.1f}\")\n    print(f\"  Median: {valid_df['sentence_word_count'].median():.1f}\")\n    print(f\"  Min:    {valid_df['sentence_word_count'].min()}\")\n    print(f\"  Max:    {valid_df['sentence_word_count'].max()}\")\n    print()\n    print(f\"Character count:\")\n    print(f\"  Mean:   {valid_df['sentence_char_count'].mean():.1f}\")\n    print(f\"  Median: {valid_df['sentence_char_count'].median():.1f}\")\n    print(f\"  Min:    {valid_df['sentence_char_count'].min()}\")\n    print(f\"  Max:    {valid_df['sentence_char_count'].max()}\")\n    \n    # Build vocabulary\n    from collections import Counter\n    import re\n    \n    word_counts = Counter()\n    \n    for sentence in valid_df['sentence'].dropna():\n        words = re.findall(r\"[a-z0-9']+\", sentence.lower())\n        word_counts.update(words)\n    \n    print(f\"\\nVocabulary size: {len(word_counts):,} unique words\")\n    \n    # Sentence length distribution\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    \n    axes[0].hist(valid_df['sentence_word_count'], bins=30, edgecolor='black', alpha=0.7)\n    axes[0].set_xlabel('Word Count')\n    axes[0].set_ylabel('Frequency')\n    axes[0].set_title('Sentence Length Distribution (Words)')\n    axes[0].axvline(valid_df['sentence_word_count'].mean(), color='red', linestyle='--', label=f'Mean: {valid_df[\"sentence_word_count\"].mean():.1f}')\n    axes[0].legend()\n    \n    axes[1].hist(valid_df['sentence_char_count'], bins=30, edgecolor='black', alpha=0.7, color='orange')\n    axes[1].set_xlabel('Character Count')\n    axes[1].set_ylabel('Frequency')\n    axes[1].set_title('Sentence Length Distribution (Characters)')\n    axes[1].axvline(valid_df['sentence_char_count'].mean(), color='red', linestyle='--', label=f'Mean: {valid_df[\"sentence_char_count\"].mean():.1f}')\n    axes[1].legend()\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Top 30 most common words\n    top_words = word_counts.most_common(30)\n    words, counts = zip(*top_words)\n    \n    plt.figure(figsize=(14, 6))\n    plt.bar(words, counts, edgecolor='black', alpha=0.7)\n    plt.xlabel('Word')\n    plt.ylabel('Frequency')\n    plt.title('Top 30 Most Common Words in Sentences')\n    plt.xticks(rotation=45, ha='right')\n    plt.tight_layout()\n    plt.show()\n    \n    print(\"\\nTop 30 words:\")\n    for i, (word, count) in enumerate(top_words, 1):\n        print(f\"  {i:2}. {word:15} {count:,}\")\n    \n    # Gloss statistics\n    print(\"GLOSS STATISTICS\")\n    print(\"-\" * 40)\n    print(f\"Gloss count per sample:\")\n    print(f\"  Mean:   {valid_df['gloss_count'].mean():.1f}\")\n    print(f\"  Median: {valid_df['gloss_count'].median():.1f}\")\n    print(f\"  Min:    {valid_df['gloss_count'].min()}\")\n    print(f\"  Max:    {valid_df['gloss_count'].max()}\")\n    \n    # Build gloss vocabulary\n    all_glosses = []\n    for gloss_list in valid_df['gloss_list']:\n        all_glosses.extend(gloss_list)\n    \n    gloss_counts = Counter(all_glosses)\n    print(f\"\\nUnique glosses: {len(gloss_counts):,}\")\n    print(f\"Total gloss occurrences: {len(all_glosses):,}\")\n    \n    # Fingerspelling analysis\n    fingerspell_glosses = [g for g in all_glosses if is_fingerspelling(g)]\n    print(f\"\\nFingerspelling:\")\n    print(f\"  Total fingerspelled letters: {len(fingerspell_glosses):,}\")\n    print(f\"  Percentage of all glosses: {100*len(fingerspell_glosses)/len(all_glosses):.1f}%\")\n    print(f\"  Samples with fingerspelling: {(valid_df['fingerspell_count'] > 0).sum():,} ({100*(valid_df['fingerspell_count'] > 0).mean():.1f}%)\")\n    \n    # Gloss distribution\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    \n    axes[0].hist(valid_df['gloss_count'], bins=30, edgecolor='black', alpha=0.7, color='green')\n    axes[0].set_xlabel('Gloss Count')\n    axes[0].set_ylabel('Frequency')\n    axes[0].set_title('Gloss Count Distribution per Sample')\n    axes[0].axvline(valid_df['gloss_count'].mean(), color='red', linestyle='--', label=f'Mean: {valid_df[\"gloss_count\"].mean():.1f}')\n    axes[0].legend()\n    \n    # Top glosses (excluding single letters)\n    non_fingerspell = {g: c for g, c in gloss_counts.items() if not is_fingerspelling(g)}\n    top_glosses = sorted(non_fingerspell.items(), key=lambda x: x[1], reverse=True)[:20]\n    glosses, counts = zip(*top_glosses)\n    \n    axes[1].barh(list(reversed(glosses)), list(reversed(counts)), edgecolor='black', alpha=0.7, color='green')\n    axes[1].set_xlabel('Frequency')\n    axes[1].set_ylabel('Gloss')\n    axes[1].set_title('Top 20 Most Common Glosses (excluding fingerspelling)')\n    \n    plt.tight_layout()\n    plt.show()\n\nif NOTEBOOK_EXPERIMENT_MODE:\n    !pip uninstall kiseki -y \n    !pip install git+https://github.com/1997MarsRover/kiseki.git\n    from kiseki import visualize\n    \n    sample_id = 1000648\n    \n    sample_path = Path(f\"/kaggle/input/motion-s-hierarchical-text-to-motion-generation-for-sign-language/Motion-Features/{1000648}.npy\")\n    assert sample_path.exists(), f\"Sample {sample_id} not found\"\n    # Focus on hands with front view\n    visualize(sample_path,\n              focus_joints='both_hands',\n              fixed_view='front',\n              fps=60, display=True)","metadata":{"_cell_guid":"dc570788-b825-4db0-ac6e-5bcee6f7a93f","_uuid":"6cde1090-2249-464d-b13b-ddd1ce0dbbac","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T07:36:59.633177Z","iopub.execute_input":"2026-03-04T07:36:59.633465Z","iopub.status.idle":"2026-03-04T07:36:59.663667Z","shell.execute_reply.started":"2026-03-04T07:36:59.633430Z","shell.execute_reply":"2026-03-04T07:36:59.662915Z"},"jupyter":{"outputs_hidden":false},"trusted":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Dataset Preprocessing","metadata":{"_cell_guid":"41c59589-8483-4509-8039-b2c9b4ae7882","_uuid":"52ebac1b-1afb-48c9-a508-5dfedbf926cd","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import pandas as pd\n\ndef preprocess_text_to_input_text(df: pd.DataFrame):\n    \"\"\"Create input_text column from gloss column.\"\"\"\n    df[\"input_text\"] = df[\"gloss\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T10:40:31.645568Z","iopub.execute_input":"2026-03-04T10:40:31.646068Z","iopub.status.idle":"2026-03-04T10:40:31.649477Z","shell.execute_reply.started":"2026-03-04T10:40:31.646042Z","shell.execute_reply":"2026-03-04T10:40:31.648848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV)\n\npreprocess_text_to_input_text(train_df)\n\nprint(\"Train:\", train_df.shape)\ntrain_df.head(1)","metadata":{"_cell_guid":"d1b8e143-fc6a-4a50-9320-401c7e80d2d9","_uuid":"7ec7ab81-511a-47ff-a173-cf2d91916e5b","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T10:40:31.983338Z","iopub.execute_input":"2026-03-04T10:40:31.984039Z","iopub.status.idle":"2026-03-04T10:40:32.365043Z","shell.execute_reply.started":"2026-03-04T10:40:31.984014Z","shell.execute_reply":"2026-03-04T10:40:32.364433Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Token Parsing and Text Tokenization\n\nConvert Token Strings to Lists of Integers  `\"45 127 88 234\" -> [...]`","metadata":{"_cell_guid":"d2e7c6b4-b155-482c-9821-3f3842334151","_uuid":"00f02de3-c018-4fc0-8483-9f6b6266f6b0","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"TOKEN_COLS=['base_tokens','residual_1','residual_2','residual_3','residual_4','residual_5']\n\nmax_token_seq_len = train_df[TOKEN_COLS].map(lambda x: 0 if pd.isna(x) else len(x.split())).max().max()\nprint(\"max_token_seq_len:\", max_token_seq_len)\n\nPAD_INDEX = 512  # same as ignore_index\nMAX_POS_LEN = 1024\n\n# assert max_token_seq_len <= MAX_POS_LEN\n\ndef parse_tokens(s: str, seq_len=MAX_POS_LEN):\n    if seq_len is None:\n        seq_len = max_token_seq_len\n\n    if pd.isna(s):\n        return np.full(seq_len, PAD_INDEX, dtype=np.int64)\n\n    arr = np.fromiter(map(int, s.split()), dtype=np.int64)\n\n    if len(arr) > seq_len:\n        arr = arr[:seq_len]\n    elif len(arr) < seq_len:\n        arr = np.pad(\n            arr,\n            (0, seq_len - len(arr)),\n            mode=\"constant\",\n            constant_values=PAD_INDEX\n        )\n\n    return arr","metadata":{"_cell_guid":"57a3981a-731c-496d-906d-314111f4115d","_uuid":"292c6f92-694b-4068-a98b-846cbfbe45b7","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T10:40:34.506500Z","iopub.execute_input":"2026-03-04T10:40:34.507173Z","iopub.status.idle":"2026-03-04T10:40:34.860830Z","shell.execute_reply.started":"2026-03-04T10:40:34.507143Z","shell.execute_reply":"2026-03-04T10:40:34.860215Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import AutoTokenizer\n\ntokenizer = AutoTokenizer.from_pretrained(T5_PATH)","metadata":{"_cell_guid":"80dcc2e4-9876-43fa-b8d0-1c727a471188","_uuid":"73a34b42-1be0-4e1d-9087-fb6c0c158c37","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T10:40:37.378772Z","iopub.execute_input":"2026-03-04T10:40:37.379370Z","iopub.status.idle":"2026-03-04T10:40:37.541765Z","shell.execute_reply.started":"2026-03-04T10:40:37.379342Z","shell.execute_reply":"2026-03-04T10:40:37.540661Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Dataset Loader","metadata":{"_cell_guid":"3e952792-7a57-456a-9685-315e314cbae3","_uuid":"21869f1a-4a8e-4131-a98f-95d51a5c0148","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class MotionDataset(Dataset):\n    def __init__(self, df, tokenizer):\n        df = df.reset_index(drop=True)\n        enc = tokenizer(\n            df[\"input_text\"].tolist(),\n            padding=\"max_length\", truncation=True,\n            max_length=128, return_tensors=\"pt\"\n        )\n        self.input_ids      = enc[\"input_ids\"]\n        self.attention_mask = enc[\"attention_mask\"]\n\n        layers = [np.stack([parse_tokens(v) for v in df[col]]) for col in TOKEN_COLS]\n        self.motion = torch.tensor(np.stack(layers), dtype=torch.long).permute(1, 2, 0)\n\n    def __len__(self): return len(self.input_ids)\n\n    def __getitem__(self, idx):\n        return {\n            \"input_ids\":      self.input_ids[idx],\n            \"attention_mask\": self.attention_mask[idx],\n            \"motion\":         self.motion[idx],\n        }","metadata":{"_cell_guid":"6e279bdb-fa91-4bad-ab22-198bdef02893","_uuid":"cdb14830-540d-4ed9-9077-7f18e14e5e68","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T10:40:40.807319Z","iopub.execute_input":"2026-03-04T10:40:40.808029Z","iopub.status.idle":"2026-03-04T10:40:40.813355Z","shell.execute_reply.started":"2026-03-04T10:40:40.808002Z","shell.execute_reply":"2026-03-04T10:40:40.812758Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split\n\ndataset = MotionDataset(train_df, tokenizer)\n\nBATCH_SIZE = 16\nTRAIN_VAL_SPLIT_RATIO = 0.5\n\ntrain_size = int(TRAIN_VAL_SPLIT_RATIO * len(dataset))\nval_size = len(dataset) - train_size\n\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\ndata_loader_args = dict(\n    batch_size=BATCH_SIZE, \n    num_workers=4,      # parallel data loading\n    pin_memory=True,    # faster CPU→GPU transfer\n    prefetch_factor=1,   # prefetch batches before GPU needs them\n)\n\ntrain_loader = DataLoader(train_dataset, shuffle=True, **data_loader_args)\nval_loader = DataLoader(val_dataset, shuffle=False, **data_loader_args)\n\nprint(\"Train size:\", train_size)\nprint(\"Val size:  \", val_size)","metadata":{"_cell_guid":"9de63461-6348-4d7e-a72d-0d1e6e380244","_uuid":"592188b2-6038-418f-ae3a-d9c6d9ebc88a","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T10:40:44.454553Z","iopub.execute_input":"2026-03-04T10:40:44.455230Z","iopub.status.idle":"2026-03-04T10:40:49.540546Z","shell.execute_reply.started":"2026-03-04T10:40:44.455201Z","shell.execute_reply":"2026-03-04T10:40:49.539681Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!nvidia-smi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T10:40:49.542067Z","iopub.execute_input":"2026-03-04T10:40:49.542532Z","iopub.status.idle":"2026-03-04T10:40:49.794955Z","shell.execute_reply.started":"2026-03-04T10:40:49.542508Z","shell.execute_reply":"2026-03-04T10:40:49.794177Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Model Definition","metadata":{"_cell_guid":"d54204ca-855b-4b46-8934-7dd9735d6615","_uuid":"ee8fa4c1-c4fb-4ed4-be2b-dd3f2e03e949","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom transformers import AutoModel, T5EncoderModel\n\n\nclass TextToMotionModel(nn.Module):\n\n    def __init__(\n        self,\n        text_base_model_name=T5_PATH,\n        hidden_dim=768,\n        max_motion_len=MAX_POS_LEN,\n        num_rvq_layers=6,\n        num_decoder_layers=4,\n        nheads=8,\n        vocab_size=513,   # tokens 0-511, PAD=512 handled by ignore_index\n        ff_dim=2048,\n        dropout=0.1,\n    ):\n        super().__init__()\n\n        self.hidden_dim = hidden_dim\n        self.vocab_size = vocab_size\n        self.max_motion_len = max_motion_len\n        self.num_rvq_layers = num_rvq_layers\n\n        # -----------------------------\n        # Text Encoder\n        # -----------------------------\n        self.text_encoder = self._build_text_encoder(text_base_model_name)\n        text_hidden = (\n            self.text_encoder.config.d_model          # T5\n            if hasattr(self.text_encoder.config, \"d_model\")\n            else self.text_encoder.config.hidden_size  # BERT/DeBERTa\n        )\n\n        # -----------------------------\n        # Projection Layer\n        # -----------------------------\n        self.text_proj = nn.Linear(text_hidden, hidden_dim)\n\n        # -----------------------------\n        # Motion Queries\n        # -----------------------------\n        self.motion_pos_emb = nn.Embedding(max_motion_len, hidden_dim)\n\n        # -----------------------------\n        # Transformer Decoder\n        # -----------------------------\n        decoder_layer = nn.TransformerDecoderLayer(\n            d_model=hidden_dim,\n            nhead=nheads,\n            dim_feedforward=ff_dim,\n            batch_first=True,\n            dropout=dropout,\n        )\n\n        self.motion_decoder = nn.TransformerDecoder(\n            decoder_layer,\n            num_layers=num_decoder_layers,\n        )\n\n        # -----------------------------\n        # RVQ Output Heads\n        # -----------------------------\n        self.output_heads = nn.ModuleList([\n            nn.Linear(hidden_dim, vocab_size)\n            for _ in range(num_rvq_layers)\n        ])\n\n    # ----------------------------------------------------\n    # Builder Functions\n    # ----------------------------------------------------\n\n    def _build_text_encoder(self, model_name):\n        \"\"\"Create text encoder backbone.\"\"\"\n        if \"t5\" in model_name.lower():\n            return T5EncoderModel.from_pretrained(model_name)\n        return AutoModel.from_pretrained(model_name)\n\n    # ----------------------------------------------------\n    # Forward Pass\n    # ----------------------------------------------------\n\n    def forward(self, input_ids, attention_mask, target_len):\n        \"\"\"\n        Args:\n            input_ids:      [B, T_text]\n            attention_mask: [B, T_text]\n            target_len:     int\n\n        Returns:\n            logits: [B, L_motion, num_rvq_layers, vocab_size]\n        \"\"\"\n\n        B = input_ids.size(0)\n\n        # -----------------------------\n        # Text Encoding\n        # -----------------------------\n        text_outputs = self.text_encoder(\n            input_ids=input_ids,\n            attention_mask=attention_mask,\n        )\n\n        text_memory = self.text_proj(\n            text_outputs.last_hidden_state\n        )  # [B, T_text, hidden_dim]\n\n        # -----------------------------\n        # Motion Query Generation\n        # -----------------------------\n        positions = (\n            torch.arange(target_len, device=input_ids.device)\n            .clamp(max=self.max_motion_len - 1)\n            .unsqueeze(0)\n            .expand(B, -1)\n        )\n\n        motion_queries = self.motion_pos_emb(positions)\n\n        # -----------------------------\n        # Cross Attention\n        # -----------------------------\n        key_padding_mask = attention_mask == 0\n\n        motion_features = self.motion_decoder(\n            tgt=motion_queries,\n            memory=text_memory,\n            memory_key_padding_mask=key_padding_mask,\n        )\n\n        # -----------------------------\n        # RVQ Prediction Heads\n        # -----------------------------\n        logits = torch.stack(\n            [head(motion_features) for head in self.output_heads],\n            dim=2,\n        )\n\n        return logits\n        \n    def summary(self):\n        total   = sum(p.numel() for p in self.parameters())\n        trained = sum(p.numel() for p in self.parameters() if p.requires_grad)\n        frozen  = total - trained\n    \n        encoder_params = sum(p.numel() for p in self.text_encoder.parameters())\n        decoder_params = sum(p.numel() for p in self.motion_decoder.parameters())\n        heads_params   = sum(p.numel() for p in self.output_heads.parameters())\n        proj_params    = sum(p.numel() for p in self.text_proj.parameters())\n        pos_params     = sum(p.numel() for p in self.motion_pos_emb.parameters())\n    \n        train_pct = 100 * trained / total if total > 0 else 0\n    \n        model_name = getattr(self.text_encoder.config, \"name_or_path\", \"unknown\")\n    \n        print(\"=\" * 50)\n        print(\"TextToMotionModel Summary\")\n        print(\"=\" * 50)\n    \n        print(f\"  Text Encoder      : {model_name}\")\n        print(f\"  Hidden Dim        : {self.hidden_dim}\")\n        print(f\"  Max Motion Len    : {self.max_motion_len}\")\n        print(f\"  RVQ Layers        : {self.num_rvq_layers}\")\n        print(f\"  Vocab Size        : {self.vocab_size}\")\n    \n        print(\"-\" * 50)\n    \n        print(f\"  Encoder Params    : {encoder_params:>12,}\")\n        print(f\"  Projection Params : {proj_params:>12,}\")\n        print(f\"  Pos Emb Params    : {pos_params:>12,}\")\n        print(f\"  Decoder Params    : {decoder_params:>12,}\")\n        print(f\"  Output Head Params: {heads_params:>12,}\")\n    \n        print(\"-\" * 50)\n    \n        print(f\"  Total Params      : {total:>12,}\")\n        print(f\"  Trainable         : {trained:>12,} ({train_pct:.2f}%)\")\n        print(f\"  Frozen            : {frozen:>12,}\")\n    \n        print(\"=\" * 50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T10:41:01.698170Z","iopub.execute_input":"2026-03-04T10:41:01.699072Z","iopub.status.idle":"2026-03-04T10:41:01.716112Z","shell.execute_reply.started":"2026-03-04T10:41:01.699035Z","shell.execute_reply":"2026-03-04T10:41:01.715513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n# from transformers import AutoModel\n# from transformers import AutoModel, T5EncoderModel\n\n# class TextToMotionModel(nn.Module):\n#     def __init__(self, \n#                  text_base_model_name=T5_PATH,\n#                  hidden_dim=512,\n#                  max_motion_len=MAX_POS_LEN,\n#                  num_layers=6,\n#                  num_decoder_layers=4,\n#                  nheads=8,\n#                  vocab_size=512):\n        \n#         super().__init__()\n        \n#         # Text encoder\n#         if \"t5\" in text_base_model_name.lower():\n#             self.text_encoder = T5EncoderModel.from_pretrained(text_base_model_name)\n#         else:\n#             self.text_encoder = AutoModel.from_pretrained(text_base_model_name)\n#         text_hidden = self.text_encoder.config.hidden_size\n        \n#         # Project full text sequence to motion hidden space\n#         self.text_proj = nn.Linear(text_hidden, hidden_dim)\n        \n#         # Motion positional embeddings\n#         self.motion_pos_emb = nn.Embedding(max_motion_len, hidden_dim)\n        \n#         # TransformerDecoder cross-attends into text\n#         decoder_layer = nn.TransformerDecoderLayer(\n#             d_model=hidden_dim,\n#             nhead=nheads,\n#             dim_feedforward=2048,\n#             batch_first=True,\n#             dropout=0.1\n#         )\n#         self.motion_decoder = nn.TransformerDecoder(\n#             decoder_layer,\n#             num_layers=6,\n#         )\n        \n#         # Separate head per RVQ layer\n#         # Respects RVQ hierarchy — each layer can specialize\n#         self.output_heads = nn.ModuleList([\n#             nn.Linear(hidden_dim, vocab_size) for _ in range(num_layers)\n#         ])\n        \n#         self.num_layers = num_layers\n#         self.vocab_size = vocab_size\n#         self.max_motion_len = max_motion_len\n\n#     def forward(self, input_ids, attention_mask, target_len):\n#         \"\"\"\n#         input_ids:      [B, T_text]\n#         attention_mask: [B, T_text]\n#         target_len:     int (motion sequence length to generate)\n#         returns logits: [B, target_len, 6, 512]\n#         \"\"\"\n#         B = input_ids.size(0)\n\n#         # Encode full text sequence\n#         text_outputs = self.text_encoder(\n#             input_ids=input_ids,\n#             attention_mask=attention_mask\n#         )\n#         text_memory = self.text_proj(\n#             text_outputs.last_hidden_state   # [B, T_text, hidden_dim]\n#         )\n\n#         # Motion query positions\n#         positions = (\n#             torch.arange(target_len, device=input_ids.device)\n#             .clamp(max=self.max_motion_len - 1)\n#             .unsqueeze(0)\n#             .expand(B, -1)\n#         )\n#         motion_queries = self.motion_pos_emb(positions)  # [B, L, hidden_dim]\n\n#         # Each motion frame learns which gloss tokens are relevant to it\n#         # key_padding_mask: True = ignore this position (opposite of attention_mask)\n#         key_padding_mask = (attention_mask == 0)         # [B, T_text]\n\n#         motion_features = self.motion_decoder(\n#             tgt=motion_queries,                          # [B, L, hidden_dim]\n#             memory=text_memory,                          # [B, T_text, hidden_dim]\n#             memory_key_padding_mask=key_padding_mask\n#         )  # [B, L, hidden_dim]\n\n#         # Separate head per RVQ layer → stack into [B, L, 6, 512]\n#         logits = torch.stack(\n#             [head(motion_features) for head in self.output_heads],\n#             dim=2\n#         )  # [B, L, 6, vocab_size]\n\n#         return logits","metadata":{"_cell_guid":"b2b7f2df-7896-4efd-b854-709b58036986","_uuid":"f5a7bf64-745d-4924-b00b-6a9a257e0b39","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T07:40:29.252122Z","iopub.execute_input":"2026-03-04T07:40:29.252337Z","iopub.status.idle":"2026-03-04T07:40:29.262369Z","shell.execute_reply.started":"2026-03-04T07:40:29.252316Z","shell.execute_reply":"2026-03-04T07:40:29.261642Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = TextToMotionModel()\nmodel.to(device)\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2026-03-04T10:41:06.970736Z","iopub.execute_input":"2026-03-04T10:41:06.971492Z","iopub.status.idle":"2026-03-04T10:41:07.330911Z","shell.execute_reply.started":"2026-03-04T10:41:06.971463Z","shell.execute_reply":"2026-03-04T10:41:07.330266Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import subprocess\n# tokenizer.save_pretrained(\"/kaggle/working/t5-base-local/\")\n# model.text_encoder.save_pretrained(\"/kaggle/working/t5-base-local/\")\n# print(\"Saved!\")\n# %%bash\n# apt-get install -qq pv\n# tar -cf - -C /kaggle/working/ t5-base-local/ | pv | gzip > /kaggle/working/t5-base-local.tar.gz\n# echo \"Done!\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T07:33:30.868308Z","iopub.execute_input":"2026-03-04T07:33:30.868564Z","iopub.status.idle":"2026-03-04T07:33:30.872240Z","shell.execute_reply.started":"2026-03-04T07:33:30.868542Z","shell.execute_reply":"2026-03-04T07:33:30.871682Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Training Loop","metadata":{"_cell_guid":"5d883225-4a78-4e63-be5b-da98f62b83aa","_uuid":"6617a699-ac94-4cd8-842d-20ead56ffa2f","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.amp import autocast, GradScaler\nfrom transformers import get_linear_schedule_with_warmup\n\nMAX_GRAD_NORM = 0.5\nWARMUP_RATIO  = 0.05\n\nclass Trainer:\n    def __init__(self, model, train_loader, val_loader, optimizer, device=\"cuda\", patience=3):\n        self.model = model.to(device)\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.device = device\n        \n        self.optimizer = optimizer\n        self.criterion = nn.CrossEntropyLoss(ignore_index=PAD_INDEX, label_smoothing=0.1)\n        self.scaler = GradScaler()\n        self.patience = patience\n        \n        total_steps   = len(train_loader) * EPOCHS\n        warmup_steps  = int(total_steps * WARMUP_RATIO)\n        \n        self.scheduler = get_linear_schedule_with_warmup(\n            self.optimizer,\n            num_warmup_steps=warmup_steps,\n            num_training_steps=total_steps\n        )\n\n        self.train_losses = []\n        self.val_losses = []\n        self.best_val_loss = float(\"inf\")\n        self.patience_counter = 0\n\n    def validate(self, epochIdx=None):\n        self.model.eval()\n        total_loss = 0\n    \n        progress_bar = tqdm(\n            self.val_loader,\n            desc=f\"[VAL] Epoch {epochIdx+1}\" if epochIdx is not None else \"Validation\",\n            leave=False\n        )\n    \n        with torch.no_grad():\n            for batch in progress_bar:\n                ids = batch[\"input_ids\"].to(self.device)\n                mask = batch[\"attention_mask\"].to(self.device)\n                motion = batch[\"motion\"].to(self.device)\n    \n                B, L, _ = motion.shape\n                logits = self.model(ids, mask, L)\n                loss = self.criterion(\n                    logits.view(B * L * 6, self.model.vocab_size),  # ✓\n                    motion.reshape(-1)\n                )\n                \n                total_loss += loss.item()\n    \n                progress_bar.set_postfix(val_loss=loss.item())\n    \n        return total_loss / len(self.val_loader)\n\n    def train_epoch(self, epochIdx):\n        self.model.train()\n        total_loss = 0\n\n        progress_bar = tqdm(\n            self.train_loader,\n            desc=f\"[TRAIN] Epoch {epochIdx+1}\",\n            leave=False\n        )\n\n        \n        for batch in progress_bar:\n            input_ids = batch[\"input_ids\"].to(self.device)\n            attention_mask = batch[\"attention_mask\"].to(self.device)\n            targets = batch[\"motion\"].to(self.device)  # [B, L, 6]\n\n            B, L, _ = targets.shape\n\n            self.optimizer.zero_grad()\n\n            with autocast(device_type=self.device):\n                logits = self.model(\n                    input_ids=input_ids,\n                    attention_mask=attention_mask,\n                    target_len=L\n                )  # [B, L, 6, 512]\n\n                # reshape for CE\n                logits = logits.view(B * L * 6, logits.size(-1))                \n                targets = targets.view(-1)\n                loss = self.criterion(logits, targets)\n\n            self.scaler.scale(loss).backward()\n            self.scaler.unscale_(self.optimizer)\n            torch.nn.utils.clip_grad_norm_(self.model.parameters(), MAX_GRAD_NORM)\n            self.scaler.step(self.optimizer)\n            self.scaler.update()\n            self.scheduler.step()\n\n            loss_item = loss.item()\n            \n            total_loss += loss_item\n            progress_bar.set_postfix(loss=loss_item)\n\n        train_loss = total_loss / len(self.train_loader)\n        val_loss = self.validate(epochIdx)    \n\n        self.train_losses.append(train_loss)\n        self.val_losses.append(val_loss)\n\n        saved = True\n        if val_loss < self.best_val_loss:\n            self.best_val_loss = val_loss\n            self.patience_counter = 0\n            torch.save(self.model.state_dict(), BEST_MODEL_PATH)\n        else:\n            self.patience_counter += 1\n            saved = False\n            print(f\"  No improvement ({self.patience_counter}/{self.patience})\")\n            if self.patience_counter >= self.patience:\n                print(\"  Early stopping triggered!\")\n                return None, None, None  # signal to stop\n\n        torch.save({\n            \"epoch\":      epochIdx,\n            \"model\":      self.model.state_dict(),\n            \"optimizer\":  self.optimizer.state_dict(),\n            \"scheduler\":  self.scheduler.state_dict(),\n            \"train_loss\": train_loss,\n            \"val_loss\":   val_loss,\n            \"best_val\":   self.best_val_loss,\n        }, LAST_MODEL_PATH)\n        \n        return train_loss, val_loss, saved\n\n\nEPOCHS = 5\nWEIGHT_DECAY  = 0.05\nPATIENCE = 5\n\noptimizer = torch.optim.AdamW([\n    {\"params\": model.text_encoder.parameters(),   \"lr\": 2e-6},\n    {\"params\": model.text_proj.parameters(),      \"lr\": 2e-5},\n    {\"params\": model.motion_pos_emb.parameters(),\"lr\": 1e-4},\n    {\"params\": model.motion_decoder.parameters(),\"lr\": 1e-4},\n    {\"params\": model.output_heads.parameters(),  \"lr\": 1e-4},\n], weight_decay=WEIGHT_DECAY)\n\ntrainer = Trainer(model=model, train_loader=train_loader, val_loader=val_loader, optimizer=optimizer, device=device, patience=PATIENCE)\n\nprint(\"=\" * 50)\nprint(\"Training Configuration\")\nprint(\"=\" * 50)\nprint(f\"Epochs        : {EPOCHS}\")\nprint(f\"Warmup Ratio  : {WARMUP_RATIO}\")\nprint(f\"Weight Decay  : {WEIGHT_DECAY}\")\nprint(f\"Batch Size    : {BATCH_SIZE}\")\nprint(f\"Max Grad Norm : {MAX_GRAD_NORM}\")\nprint(f\"Max Pos Len   : {MAX_POS_LEN}\")\nprint(f\"Device        : {device}\")\nprint(\"=\" * 50)\nprint()\nprint()\n\nbest_epoch = 0\nfor epoch in range(EPOCHS):\n    train_loss, val_loss, saved = trainer.train_epoch(epoch)\n\n    if (train_loss, val_loss, saved) == (None, None, None):\n        print(f\"Stopped early at epoch {epoch+1}\")\n        break\n    \n    print(f\"Epoch {epoch+1:02d} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n    if saved:\n        print(f\"  New best model saved (val_loss={val_loss:.4f})\")\n        best_epoch = epoch + 1\n    else:\n        print(f\"  No improvement in validation loss ({trainer.patience_counter}/{trainer.patience})\")\n\n\nprint()\nprint(f\"Best Epoch       : {best_epoch}\")\nprint(f\"Best Val Loss    : {trainer.best_val_loss:.4f}\")","metadata":{"_cell_guid":"e84ba4bc-83d5-4783-a214-fe8f486c3920","_uuid":"996e657e-853a-4b42-ad17-74b854797a3a","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T10:41:28.369681Z","iopub.execute_input":"2026-03-04T10:41:28.370049Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.figure(figsize=(8,5))\nplt.plot(trainer.train_losses, label=\"Train Loss\")\nplt.plot(trainer.val_losses, label=\"Validation Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs Validation Loss\")\nplt.legend()\nplt.grid(True)\nplt.savefig(\"/kaggle/working/loss_curve.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()","metadata":{"_cell_guid":"7ef94f28-b371-4ab4-9076-4ec93f94d88c","_uuid":"d5ea75d0-5971-4e71-9bc8-76eb52d4b5aa","collapsed":false,"execution":{"iopub.status.busy":"2026-03-04T07:43:30.121330Z","iopub.status.idle":"2026-03-04T07:43:30.121608Z","shell.execute_reply.started":"2026-03-04T07:43:30.121486Z","shell.execute_reply":"2026-03-04T07:43:30.121504Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. Test Dataset","metadata":{"_cell_guid":"180aa6ba-859d-427f-961b-57cc37fff674","_uuid":"72d47a39-9983-4dbb-947c-86610e873ecb","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"test_df = pd.read_csv(TEST_CSV)\n\npreprocess_text_to_input_text(test_df)\n\nprint(\"Test:\", test_df.shape)","metadata":{"execution":{"iopub.execute_input":"2026-03-03T19:48:53.804009Z","iopub.status.busy":"2026-03-03T19:48:53.803597Z","iopub.status.idle":"2026-03-03T19:48:53.828917Z","shell.execute_reply":"2026-03-03T19:48:53.828206Z","shell.execute_reply.started":"2026-03-03T19:48:53.803966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, df, tokenizer):\n        df = df.reset_index(drop=True)\n        enc = tokenizer(\n            df[\"input_text\"].astype(str).tolist(),\n            padding=\"max_length\", truncation=True,\n            max_length=128, return_tensors=\"pt\"\n        )\n        self.ids            = df[\"id\"].values\n        self.input_ids      = enc[\"input_ids\"]\n        self.attention_mask = enc[\"attention_mask\"]\n\n    def __len__(self): return len(self.ids)\n\n    def __getitem__(self, idx):\n        return {\n            \"id\":             self.ids[idx],\n            \"input_ids\":      self.input_ids[idx],\n            \"attention_mask\": self.attention_mask[idx],\n        }","metadata":{"_cell_guid":"c6edc2c5-728a-4051-8079-bdc6039c700b","_uuid":"f345cca9-8152-46af-9092-c4d25f80424e","collapsed":false,"execution":{"iopub.execute_input":"2026-03-03T19:48:56.667219Z","iopub.status.busy":"2026-03-03T19:48:56.666371Z","iopub.status.idle":"2026-03-03T19:48:56.673564Z","shell.execute_reply":"2026-03-03T19:48:56.672708Z","shell.execute_reply.started":"2026-03-03T19:48:56.667184Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split\n\ndataset = TestDataset(test_df, tokenizer)\n\nBATCH_SIZE = 16\n\ntest_loader = DataLoader(\n    dataset,\n    shuffle=False,\n    batch_size=BATCH_SIZE, \n    num_workers=4,      # parallel data loading\n    pin_memory=True,    # faster CPU→GPU transfer\n    prefetch_factor=1,   # prefetch batches before GPU needs them\n)","metadata":{"execution":{"iopub.execute_input":"2026-03-03T19:49:01.980735Z","iopub.status.busy":"2026-03-03T19:49:01.980378Z","iopub.status.idle":"2026-03-03T19:49:02.401917Z","shell.execute_reply":"2026-03-03T19:49:02.401226Z","shell.execute_reply.started":"2026-03-03T19:49:01.980702Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. Inference","metadata":{"_cell_guid":"b8a17f8d-d82f-4ad5-8b80-f8fdc8cef532","_uuid":"34273310-5a8f-4a0f-bc9c-48e30796767d","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"outputs = []\n\nmodel.eval()\n\nwith torch.no_grad():\n    for batch in tqdm(test_loader):\n        ids  = batch[\"input_ids\"].to(device)\n        mask = batch[\"attention_mask\"].to(device)\n\n        logits = model(ids, mask, 800)           # FIX: 120 → 800 (submission requires 40-800)\n        preds  = logits.argmax(-1).cpu().numpy() # [B, 800, 6]\n\n        for i in range(len(batch[\"id\"])):\n            sid = int(batch[\"id\"][i])\n            row = [sid]\n            for layer_idx in range(6):           # FIX: index each layer separately\n                row.append(\" \".join(map(str, preds[i, :, layer_idx])))\n            outputs.append(row)\n","metadata":{"_cell_guid":"afc29ca2-5ee8-44f4-ad69-d8f35706fb89","_uuid":"8ad03b88-a119-4966-a035-ec46693a1584","collapsed":false,"execution":{"iopub.execute_input":"2026-03-03T19:49:06.944509Z","iopub.status.busy":"2026-03-03T19:49:06.944193Z","iopub.status.idle":"2026-03-03T19:49:46.965283Z","shell.execute_reply":"2026-03-03T19:49:46.964417Z","shell.execute_reply.started":"2026-03-03T19:49:06.944481Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 9. Submission","metadata":{"_cell_guid":"3b3d7096-f162-4e36-a2e8-b03343c665e2","_uuid":"a766074b-7267-41f0-9245-bf0ae33bc3eb","collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"test_ids = test_df[\"id\"].values\nsubmission = pd.DataFrame(outputs, columns=[\"id\"] + TOKEN_COLS)\nsubmission = submission.set_index(\"id\").loc[test_ids].reset_index()\nsubmission.to_csv(\"/kaggle/working/submission.csv\", index=False)","metadata":{"execution":{"iopub.execute_input":"2026-03-03T19:49:55.600094Z","iopub.status.busy":"2026-03-03T19:49:55.599744Z","iopub.status.idle":"2026-03-03T19:49:57.169849Z","shell.execute_reply":"2026-03-03T19:49:57.169122Z","shell.execute_reply.started":"2026-03-03T19:49:55.600056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"assert len(submission) == len(test_ids), \"Row count mismatch\"\nfor col in TOKEN_COLS:\n    lengths = submission[col].apply(lambda x: len(x.split()))\n    maxvals = submission[col].apply(lambda x: max(map(int, x.split())))\n    assert lengths.between(40, 800).all(), f\"{col} has invalid lengths\"\n    assert maxvals.max() <= 511, f\"{col} has tokens > 511\"\n\nprint(f\"Rows: {len(submission)} / {len(test_ids)}\")\nprint(f\"ID order correct: {(submission['id'].values == test_ids).all()}\")\nprint(\"Submission valid!\")","metadata":{"_cell_guid":"3f3e62a2-283f-48c3-9888-b6070f1e5c35","_uuid":"ef7936d5-7c29-4cbb-8020-c28f4a1bc111","collapsed":false,"execution":{"iopub.execute_input":"2026-03-03T19:50:02.100926Z","iopub.status.busy":"2026-03-03T19:50:02.100152Z","iopub.status.idle":"2026-03-03T19:50:04.959787Z","shell.execute_reply":"2026-03-03T19:50:04.958993Z","shell.execute_reply.started":"2026-03-03T19:50:02.100896Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"DONE!\")","metadata":{"execution":{"iopub.execute_input":"2026-03-03T19:50:04.961655Z","iopub.status.busy":"2026-03-03T19:50:04.961131Z","iopub.status.idle":"2026-03-03T19:50:04.966316Z","shell.execute_reply":"2026-03-03T19:50:04.965506Z","shell.execute_reply.started":"2026-03-03T19:50:04.961625Z"}},"outputs":[],"execution_count":null}]}