{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceType":"competition","sourceId":19018,"databundleVersionId":2703900},{"sourceType":"datasetVersion","sourceId":1062669,"datasetId":588377,"databundleVersionId":1092212}],"dockerImageVersionId":31331,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.environ['XLA_USE_BF16'] = '1'\n\nimport numpy as np \nimport pandas as pd \nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:44:50.069700Z","iopub.execute_input":"2026-03-28T17:44:50.070000Z","iopub.status.idle":"2026-03-28T17:44:50.422130Z","shell.execute_reply.started":"2026-03-28T17:44:50.069980Z","shell.execute_reply":"2026-03-28T17:44:50.421096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch_xla","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:44:50.422857Z","iopub.execute_input":"2026-03-28T17:44:50.423117Z","iopub.status.idle":"2026-03-28T17:45:07.059500Z","shell.execute_reply.started":"2026-03-28T17:44:50.423097Z","shell.execute_reply":"2026-03-28T17:45:07.058441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train1 = pd.read_csv('/kaggle/input/competitions/jigsaw-multilingual-toxic-comment-classification/jigsaw-toxic-comment-train.csv')\ntrain2 = pd.read_csv('/kaggle/input/competitions/jigsaw-multilingual-toxic-comment-classification/jigsaw-unintended-bias-train.csv')\ntrain2['toxic'] = train2['toxic'].round().astype(int)\nvalid_df = pd.read_csv('/kaggle/input/competitions/jigsaw-multilingual-toxic-comment-classification/validation.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:45:07.060133Z","iopub.execute_input":"2026-03-28T17:45:07.060508Z","iopub.status.idle":"2026-03-28T17:45:27.624910Z","shell.execute_reply.started":"2026-03-28T17:45:07.060489Z","shell.execute_reply":"2026-03-28T17:45:27.623663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Loading Google API Translations...\")\nlangs = ['es', 'fr', 'it', 'pt', 'ru', 'tr']\ntranslated_data = []\n\nbase_path = '/kaggle/input/datasets/miklgr500/jigsaw-train-multilingual-coments-google-api' \n\nfor lang in langs:\n    df = pd.read_csv(f'{base_path}/jigsaw-toxic-comment-train-google-{lang}-cleaned.csv')\n    translated_data.append(df[['comment_text', 'toxic']])\n\ngoogle_df = pd.concat(translated_data)\n\n# 3. Smart Subsampling (The TPU Saver)\n# We take ALL the toxic foreign comments, but only 50,000 random clean ones.\ngoogle_toxic = google_df[google_df['toxic'] == 1]\ngoogle_clean = google_df[google_df['toxic'] == 0].sample(n=50000, random_state=42)\n\n# 4. Build the Master Stage 1 Dataset\ntrain_df = pd.concat([\n    train1[['comment_text', 'toxic']],\n    train2[['comment_text', 'toxic']].query('toxic==1'),\n    train2[['comment_text', 'toxic']].query('toxic==0').sample(n=100000, random_state=42),\n    google_toxic,\n    google_clean\n]).sample(frac=1, random_state=42).reset_index(drop=True)\n\nprint(f\"Final training rows: {len(train_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:45:27.625508Z","iopub.execute_input":"2026-03-28T17:45:27.625687Z","iopub.status.idle":"2026-03-28T17:45:44.117015Z","shell.execute_reply.started":"2026-03-28T17:45:27.625669Z","shell.execute_reply":"2026-03-28T17:45:44.115848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import AutoTokenizer, AutoModelForSequenceClassification \n\ntokenizer = AutoTokenizer.from_pretrained(\"xlm-roberta-large\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:45:44.117785Z","iopub.execute_input":"2026-03-28T17:45:44.117978Z","iopub.status.idle":"2026-03-28T17:45:51.204488Z","shell.execute_reply.started":"2026-03-28T17:45:44.117959Z","shell.execute_reply":"2026-03-28T17:45:51.203326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset\n\nclass ToxicDataset(Dataset):\n\n    def __init__(self, texts, tokenizer,max_length, labels=None):\n        self.texts = texts\n        self.labels = labels\n        self.tokenizer = tokenizer\n        self.max_length = max_length\n\n    def __len__(self):\n        return len(self.texts)\n\n    def __getitem__(self, idx):\n        enc = self.tokenizer(\n            self.texts[idx],\n            max_length = self.max_length,\n            padding = 'max_length',\n            truncation = True,\n            return_tensors = 'pt'\n        )\n        item = {\n            'input_ids' : enc['input_ids'].squeeze(),\n            'attention_mask': enc['attention_mask'].squeeze()\n        }\n\n        if self.labels is not None:\n            item['labels'] = torch.tensor(self.labels[idx], dtype = torch.float)\n        return item","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:45:51.205218Z","iopub.execute_input":"2026-03-28T17:45:51.205710Z","iopub.status.idle":"2026-03-28T17:45:51.210439Z","shell.execute_reply.started":"2026-03-28T17:45:51.205692Z","shell.execute_reply":"2026-03-28T17:45:51.209648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = ToxicDataset(\n    texts=train_df['comment_text'].tolist(),\n    tokenizer=tokenizer,\n    max_length=192,\n    labels=train_df['toxic'].tolist()\n)\n\nvalid_dataset = ToxicDataset(\n    texts=valid_df['comment_text'].tolist(),\n    tokenizer=tokenizer,\n    max_length=192,\n    labels=valid_df['toxic'].tolist()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:45:51.210936Z","iopub.execute_input":"2026-03-28T17:45:51.211104Z","iopub.status.idle":"2026-03-28T17:45:51.490704Z","shell.execute_reply.started":"2026-03-28T17:45:51.211088Z","shell.execute_reply":"2026-03-28T17:45:51.489701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.runtime as xr\nimport torch_xla.distributed.xla_multiprocessing as xmp\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data.distributed import DistributedSampler\nfrom torch.optim import AdamW\nimport torch_xla.runtime as xr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T17:41:27.211401Z","iopub.execute_input":"2026-03-28T17:41:27.211671Z","iopub.status.idle":"2026-03-28T17:41:27.215095Z","shell.execute_reply.started":"2026-03-28T17:41:27.211651Z","shell.execute_reply":"2026-03-28T17:41:27.214325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _mp_fn(index):\n    device = xm.xla_device()\n    \n    # --- SETUP SAMPLERS & LOADERS ---\n    train_sampler = DistributedSampler(\n        train_dataset, num_replicas=xr.world_size(), rank=xr.global_ordinal(), shuffle=True\n    )\n    valid_sampler = DistributedSampler(\n        valid_dataset, num_replicas=xr.world_size(), rank=xr.global_ordinal(), shuffle=True\n    )\n\n\n    train_loader = DataLoader(\n        train_dataset, batch_size=16, sampler=train_sampler, num_workers=4, drop_last=True\n    )\n    valid_loader = DataLoader(\n        valid_dataset, batch_size=16, sampler=valid_sampler, num_workers=4, drop_last=True\n    )\n\n    # --- MODEL & OPTIMIZER ---\n    model = AutoModelForSequenceClassification.from_pretrained('xlm-roberta-large', num_labels=1)\n    model = model.to(device)  \n\n    # Learning rate lowered to 1e-5 for the large model\n    optimizer = AdamW(model.parameters(), lr=1e-5)\n\n    # ==========================================\n    # STAGE 1: English Foundation Training\n    # ==========================================\n    xm.master_print(\"\\n=== STAGE 1: English Foundation Training (2 Epochs) ===\")\n    for epoch in range(2):\n        model.train()\n        para_loader = pl.ParallelLoader(train_loader, [device])\n        \n        xm.master_print(f\"Starting Stage 1, Epoch {epoch+1}/2...\")\n\n        for step, batch in enumerate(para_loader.per_device_loader(device)):\n            ids, masks, labels = batch[\"input_ids\"], batch[\"attention_mask\"], batch[\"labels\"]\n            \n            optimizer.zero_grad()\n            outputs = model(ids, attention_mask=masks, labels=labels)\n            loss = outputs.loss\n                \n            loss.backward()          \n            xm.optimizer_step(optimizer) \n            \n            if step % 200 == 0:\n                xm.master_print(f\"Stage 1 | Epoch {epoch+1} | Step {step} | Loss: {loss.item():.4f}\")\n    # ==========================================\n    # STAGE 2: Cross-Lingual Fine-Tuning\n    # ==========================================\n    xm.master_print(\"\\n=== STAGE 2: Cross-Lingual Fine-Tuning (1 Epoch) ===\")\n    model.train() # Model stays in training mode\n    valid_para_loader = pl.ParallelLoader(valid_loader, [device])\n    \n    for step, batch in enumerate(valid_para_loader.per_device_loader(device)):\n        ids, masks, labels = batch[\"input_ids\"], batch[\"attention_mask\"], batch[\"labels\"]\n        \n        optimizer.zero_grad()\n        outputs = model(ids, attention_mask=masks, labels=labels)\n        loss = outputs.loss\n            \n        loss.backward()          \n        xm.optimizer_step(optimizer) \n        \n        if step % 50 == 0: # Printing more often because the validation set is much smaller\n            xm.master_print(f\"Stage 2 | Fine-Tuning | Step {step} | Loss: {loss.item():.4f}\")\n\n    # Optional: Save the model after both stages are complete\n    xm.save(model.state_dict(), \"xlm_roberta_large_twostage.pth\")\nimport os\n\nif __name__ == '__main__':\n    # Clear out Kaggle's pre-configured TensorFlow variables so PyTorch XLA doesn't get confused\n    os.environ.pop('TPU_PROCESS_ADDRESSES', None)\n    os.environ.pop('CLOUD_TPU_TASK_ID', None)\n    \n    # Launch across all 8 cores\n    xmp.spawn(_mp_fn, args=(), start_method='fork')\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch_xla.core.xla_model as xm\nfrom torch.utils.data import DataLoader\nimport numpy as np\nimport pandas as pd\n\n# 1. Initialize a single TPU core for inference\ndevice = xm.xla_device()\n\n# 2. Rebuild the empty model architecture\nprint(\"Loading model architecture...\")\nmodel = AutoModelForSequenceClassification.from_pretrained('xlm-roberta-large', num_labels=1)\n\n# 3. Load your hard-earned trained weights\nprint(\"Loading saved weights...\")\nmodel.load_state_dict(torch.load(\"xlm_roberta_large_twostage.pth\"))\nmodel.to(device)\nmodel.eval()\n\n# 4. Prepare the Test Dataset (Upgraded to max_length 192)\ntest_df = pd.read_csv('/kaggle/input/competitions/jigsaw-multilingual-toxic-comment-classification/test.csv')\n\ntest_dataset = ToxicDataset(\n    texts = test_df['content'].tolist(),\n    tokenizer = tokenizer,\n    max_length = 192, # CRITICAL UPGRADE\n    labels = None\n)\n\n# Batch size dropped to 16 to handle the Large model memory footprint\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size = 16, \n    shuffle = False,\n    num_workers = 4\n)\n\n# 5. Generate Predictions\nprint(\"Starting predictions...\")\ntest_preds = []\n\nwith torch.no_grad():\n    # We can use standard tqdm here since we are only on 1 core\n    from tqdm import tqdm\n    for batch in tqdm(test_loader, desc=\"Predicting\"):\n        ids = batch[\"input_ids\"].to(device)\n        masks = batch[\"attention_mask\"].to(device)\n        \n        outputs = model(ids, attention_mask=masks)\n        preds = torch.sigmoid(outputs.logits).cpu().numpy()\n        test_preds.append(preds)\n\ntest_preds = np.vstack(test_preds).flatten()\n\n# 6. Format Submission\nsubmission = pd.read_csv('/kaggle/input/competitions/jigsaw-multilingual-toxic-comment-classification/sample_submission.csv')\nsubmission['toxic'] = test_preds\nsubmission.to_csv('submission.csv', index=False)\nprint(\"Submission saved successfully.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}