{"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":396802,"sourceType":"datasetVersion","datasetId":175990}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==============================================================================\n# PALIGEMMA FASHION - FIXED LABEL MASKING (THIS WILL WORK!)\n# ==============================================================================\n\nimport sys\nimport os\n\nprint(\"=\"*70, flush=True)\nprint(\"🚀 PALIGEMMA FASHION - PROPER LABEL MASKING FIX\", flush=True)\nprint(\"=\"*70, flush=True)\n\n!pip install -q transformers accelerate peft bitsandbytes datasets pillow scikit-learn huggingface_hub\n\nprint(\"\\n🔐 Authenticating HuggingFace...\", flush=True)\nfrom kaggle_secrets import UserSecretsClient\nfrom huggingface_hub import login\n\ntry:\n    login(token=UserSecretsClient().get_secret(\"HF_TOKEN\"))\n    print(\"✅ HuggingFace authenticated\", flush=True)\nexcept Exception as e:\n    print(f\"❌ AUTHENTICATION FAILED: {e}\", flush=True)\n    raise\n\nprint(\"✅ Setup complete\\n\", flush=True)\n\n# ==============================================================================\n# Imports\n# ==============================================================================\nimport glob\nimport time\nimport pandas as pd\nimport numpy as np\nimport torch\nimport warnings\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom dataclasses import dataclass\nfrom typing import Any, Dict, List\nfrom transformers import AutoProcessor, AutoModelForVision2Seq, TrainingArguments, Trainer, TrainerCallback\nfrom transformers.utils import logging as hf_logging\nfrom peft import LoraConfig, get_peft_model\nfrom torch.utils.data import Dataset\n\nwarnings.filterwarnings('ignore')\nhf_logging.set_verbosity_error()\n\nprint(\"✅ Imports loaded\", flush=True)\nprint(f\"🔥 PyTorch: {torch.__version__}\", flush=True)\nprint(f\"🔥 CUDA available: {torch.cuda.is_available()}\", flush=True)\nprint(f\"🔥 GPU count: {torch.cuda.device_count()}\", flush=True)\nfor i in range(torch.cuda.device_count()):\n    print(f\"   GPU {i}: {torch.cuda.get_device_name(i)}\", flush=True)\n\n# ==============================================================================\n# Load Dataset (ULTRA-OPTIMIZED - 30% sample)\n# ==============================================================================\nprint(\"\\n\" + \"=\"*70, flush=True)\nprint(\"LOADING FASHION DATASET (ULTRA-OPTIMIZED)\", flush=True)\nprint(\"=\"*70, flush=True)\n\ndataset_paths = glob.glob('/kaggle/input/*fashion*')\nif not dataset_paths:\n    raise FileNotFoundError(\"Dataset not found!\")\nDATASET_PATH = dataset_paths[0]\nprint(f\"✅ Dataset: {DATASET_PATH}\", flush=True)\n\nstyles_csv = f'{DATASET_PATH}/styles.csv'\nstyles_df = pd.read_csv(styles_csv, on_bad_lines='skip')\nprint(f\"✅ Loaded {len(styles_df)} products\", flush=True)\n\ndef create_caption(row):\n    parts = []\n    for col in ['articleType', 'baseColour', 'season', 'usage']:\n        if pd.notna(row.get(col)):\n            parts.append(str(row[col]))\n    if not parts and pd.notna(row.get('productDisplayName')):\n        return row['productDisplayName']\n    return ', '.join(parts) if parts else \"fashion product\"\n\nstyles_df['img_path'] = styles_df['id'].apply(lambda x: f'{DATASET_PATH}/images/{x}.jpg')\nstyles_df['caption'] = styles_df.apply(create_caption, axis=1)\n\nfashion_df = styles_df[['img_path', 'caption']].dropna()\nfashion_df = fashion_df[fashion_df['caption'].str.len() > 0]\n\nvalid_mask = fashion_df['img_path'].apply(os.path.exists)\nfashion_df = fashion_df[valid_mask]\n\n# ✅ ULTRA-OPTIMIZATION: Use only 30% of data\nfashion_df = fashion_df.sample(frac=0.3, random_state=42)\nprint(f\"✅ Sampled {len(fashion_df)} samples (30% for max speed)\", flush=True)\n\ntrain_df, val_df = train_test_split(fashion_df, test_size=0.1, random_state=42)\n\nprint(f\"\\n📊 Training: {len(train_df)} | Validation: {len(val_df)}\", flush=True)\nprint(\"=\"*70, flush=True)\n\n# ==============================================================================\n# Dataset Class (CRITICAL FIX - PROPER LABEL MASKING!)\n# ==============================================================================\nclass FashionDataset(Dataset):\n    def __init__(self, df, processor):\n        self.df = df.reset_index(drop=True)\n        self.processor = processor\n        print(f\"✅ Dataset created: {len(self.df)} samples\", flush=True)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        try:\n            row = self.df.iloc[idx]\n            img = Image.open(row['img_path']).convert(\"RGB\")\n            caption = str(row['caption'])[:50]\n\n            # ✅ FIX: Add <image> token to prefix!\n            prompt = \"<image>caption\"\n            \n            inputs = self.processor(\n                text=prompt,\n                images=img,\n                suffix=caption,\n                return_tensors=\"pt\",\n                padding=\"longest\",\n            )\n\n            inputs = {k: v.squeeze(0) for k, v in inputs.items()}\n            \n            # Create labels\n            labels = inputs[\"input_ids\"].clone()\n            labels[labels == self.processor.tokenizer.pad_token_id] = -100\n            \n            if \"token_type_ids\" in inputs:\n                prefix_mask = inputs[\"token_type_ids\"] == 0\n                labels[prefix_mask] = -100\n            \n            inputs[\"labels\"] = labels\n\n            return inputs\n\n        except Exception as e:\n            print(f\"Error loading sample {idx}: {e}\", flush=True)\n            dummy_inputs = self.processor(\n                text=\"<image>caption\",\n                images=Image.new('RGB', (224, 224), color='white'),\n                suffix=\"product\",\n                return_tensors=\"pt\",\n                padding=\"longest\",\n            )\n            dummy_inputs = {k: v.squeeze(0) for k, v in dummy_inputs.items()}\n            dummy_labels = dummy_inputs[\"input_ids\"].clone()\n            if \"token_type_ids\" in dummy_inputs:\n                dummy_labels[dummy_inputs[\"token_type_ids\"] == 0] = -100\n            dummy_inputs[\"labels\"] = dummy_labels\n            return dummy_inputs\n\n@dataclass\nclass DataCollator:\n    processor: Any\n\n    def __call__(self, batch: List[Dict[str, Any]]) -> Dict[str, Any]:\n        # ✅ Pad to max length in batch\n        max_len = max(b['input_ids'].shape[0] for b in batch)\n        \n        pixel_values = torch.stack([b['pixel_values'] for b in batch])\n        \n        input_ids_list = []\n        attention_mask_list = []\n        labels_list = []\n        token_type_ids_list = []\n        \n        for b in batch:\n            curr_len = b['input_ids'].shape[0]\n            pad_len = max_len - curr_len\n            \n            if pad_len > 0:\n                # Pad input_ids\n                input_ids = torch.cat([\n                    b['input_ids'],\n                    torch.full((pad_len,), self.processor.tokenizer.pad_token_id, dtype=b['input_ids'].dtype)\n                ])\n                \n                # Pad attention_mask\n                attention_mask = torch.cat([\n                    b['attention_mask'],\n                    torch.zeros(pad_len, dtype=b['attention_mask'].dtype)\n                ])\n                \n                # Pad labels\n                labels = torch.cat([\n                    b['labels'],\n                    torch.full((pad_len,), -100, dtype=b['labels'].dtype)\n                ])\n                \n                # Pad token_type_ids if present\n                if 'token_type_ids' in b:\n                    token_type_ids = torch.cat([\n                        b['token_type_ids'],\n                        torch.zeros(pad_len, dtype=b['token_type_ids'].dtype)\n                    ])\n                else:\n                    token_type_ids = torch.zeros(max_len, dtype=torch.long)\n            else:\n                input_ids = b['input_ids']\n                attention_mask = b['attention_mask']\n                labels = b['labels']\n                token_type_ids = b.get('token_type_ids', torch.zeros(curr_len, dtype=torch.long))\n            \n            input_ids_list.append(input_ids)\n            attention_mask_list.append(attention_mask)\n            labels_list.append(labels)\n            token_type_ids_list.append(token_type_ids)\n        \n        result = {\n            'pixel_values': pixel_values,\n            'input_ids': torch.stack(input_ids_list),\n            'attention_mask': torch.stack(attention_mask_list),\n            'labels': torch.stack(labels_list),\n        }\n        \n        # Add token_type_ids if available\n        if any('token_type_ids' in b for b in batch):\n            result['token_type_ids'] = torch.stack(token_type_ids_list)\n        \n        return result\n\n# ==============================================================================\n# Custom Callback\n# ==============================================================================\nclass ProgressCallback(TrainerCallback):\n    def __init__(self):\n        self.start_time = time.time()\n        self.last_log_step = 0\n\n    def on_log(self, args, state, control, logs=None, **kwargs):\n        if logs is not None and state.global_step - self.last_log_step >= 50:\n            self.last_log_step = state.global_step\n            elapsed = (time.time() - self.start_time) / 60\n            loss = logs.get('loss', 'N/A')\n            lr = logs.get('learning_rate', 'N/A')\n\n            loss_str = f\"{loss:.4f}\" if isinstance(loss, float) else str(loss)\n            lr_str = f\"{lr:.2e}\" if isinstance(lr, float) else str(lr)\n\n            print(f\"⏱️  Step {state.global_step}/{state.max_steps} | Loss: {loss_str} | LR: {lr_str} | {elapsed:.1f}min\", flush=True)\n            \n            # Check for bad loss\n            if isinstance(loss, float):\n                if loss > 5.0 and state.global_step > 100:\n                    print(f\"⚠️  WARNING: High loss ({loss:.4f}) - check data!\", flush=True)\n                if np.isnan(loss) or np.isinf(loss):\n                    print(f\"❌ STOPPING: NaN/Inf loss!\", flush=True)\n                    control.should_training_stop = True\n\n# ==============================================================================\n# Load Model\n# ==============================================================================\nprint(\"\\n\" + \"=\"*70, flush=True)\nprint(\"LOADING PALIGEMMA MODEL\", flush=True)\nprint(\"=\"*70, flush=True)\n\nmodel_id = \"google/paligemma-3b-mix-224\"\n\nprint(\"Loading processor...\", flush=True)\nprocessor = AutoProcessor.from_pretrained(model_id)\nprint(\"✅ Processor loaded\", flush=True)\n\nprint(\"Loading base model...\", flush=True)\nmodel = AutoModelForVision2Seq.from_pretrained(\n    model_id,\n    torch_dtype=torch.bfloat16,\n    device_map=\"auto\",\n)\n\nmodel.config.use_cache = False\nprint(\"✅ Model loaded\", flush=True)\n\nprint(\"Applying LoRA...\", flush=True)\nlora_config = LoraConfig(\n    r=8,\n    lora_alpha=8,\n    lora_dropout=0.05,\n    target_modules=[\"q_proj\", \"v_proj\"],\n    bias=\"none\",\n    task_type=\"CAUSAL_LM\"\n)\nmodel = get_peft_model(model, lora_config)\nmodel.enable_input_require_grads()\n\ntrainable, total = model.get_nb_trainable_parameters()\nprint(f\"✅ LoRA applied: {trainable:,}/{total:,} trainable ({trainable/total*100:.2f}%)\", flush=True)\nprint(\"=\"*70, flush=True)\n\n# ==============================================================================\n# Create Datasets\n# ==============================================================================\nprint(\"\\nCreating datasets...\", flush=True)\ntrain_ds = FashionDataset(train_df, processor)\nval_ds = FashionDataset(val_df, processor)\nprint(f\"✅ Datasets ready\", flush=True)\n\n# ==============================================================================\n# Training Config\n# ==============================================================================\ntraining_args = TrainingArguments(\n    output_dir=\"./fashion_output\",\n\n    per_device_train_batch_size=8,\n    per_device_eval_batch_size=8,\n    gradient_accumulation_steps=2,\n\n    num_train_epochs=1,\n    learning_rate=5e-5,\n    warmup_steps=50,\n    lr_scheduler_type=\"cosine\",\n    max_grad_norm=1.0,\n\n    bf16=True,\n\n    logging_steps=50,\n    logging_first_step=True,\n    eval_strategy=\"steps\",\n    eval_steps=100,\n    save_strategy=\"steps\",\n    save_steps=200,\n    save_total_limit=3,\n    load_best_model_at_end=True,\n    metric_for_best_model=\"eval_loss\",\n\n    remove_unused_columns=False,\n\n    dataloader_num_workers=2,\n    dataloader_pin_memory=True,\n\n    report_to=\"none\",\n    disable_tqdm=False,\n    log_level=\"warning\",\n\n    gradient_checkpointing=True,\n    gradient_checkpointing_kwargs={\"use_reentrant\": False},\n\n    ddp_find_unused_parameters=False,\n    weight_decay=0.01,\n)\n\nprint(\"\\n✅ TrainingArguments created\", flush=True)\n\nprint(\"Creating Trainer...\", flush=True)\ntrainer = Trainer(\n    model=model,\n    args=training_args,\n    train_dataset=train_ds,\n    eval_dataset=val_ds,\n    data_collator=DataCollator(processor),\n    callbacks=[ProgressCallback()],\n)\nprint(\"✅ Trainer created\", flush=True)\n\ngpu_count = max(1, torch.cuda.device_count())\nexpected_steps = (len(train_ds) // (training_args.per_device_train_batch_size * training_args.gradient_accumulation_steps * gpu_count)) * training_args.num_train_epochs\nprint(f\"\\n📊 Expected steps: ~{expected_steps}\", flush=True)\nprint(f\"📊 THIS SHOULD WORK NOW! ✅\", flush=True)\n\n# ==============================================================================\n# Train\n# ==============================================================================\nprint(\"\\n\" + \"=\"*70, flush=True)\nprint(\"🚀 STARTING TRAINING WITH PROPER LABEL MASKING\", flush=True)\nprint(\"=\"*70 + \"\\n\", flush=True)\n\nstart_time = time.time()\n\ntry:\n    train_result = trainer.train()\n    elapsed = (time.time() - start_time) / 3600\n\n    print(\"\\n\" + \"=\"*70, flush=True)\n    print(f\"✅ TRAINING COMPLETE! ({elapsed:.2f} hours)\", flush=True)\n    print(\"=\"*70, flush=True)\n    print(train_result, flush=True)\n\n    # Save final model\n    final_dir = \"./fashion_final\"\n    model.save_pretrained(final_dir)\n    processor.save_pretrained(final_dir)\n    print(f\"\\n💾 Model saved to: {final_dir}\", flush=True)\n\n    !zip -r fashion-model.zip {final_dir}\n    print(\"✅ Download: fashion-model.zip\", flush=True)\n\n    # Test inference\n    print(\"\\n\" + \"=\"*70, flush=True)\n    print(\"🧪 TESTING INFERENCE\", flush=True)\n    print(\"=\"*70, flush=True)\n\n    model.eval()\n\n    for i, (_, row) in enumerate(val_df.sample(3).iterrows(), 1):\n        print(f\"\\n--- Test {i}/3 ---\", flush=True)\n        img = Image.open(row['img_path']).convert(\"RGB\")\n\n        # ✅ Use same prefix as training\n        inputs = processor(\n            text=\"<image>caption\",\n            images=img,\n            return_tensors=\"pt\"\n        )\n\n        inputs = {k: v.to(model.device, dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) for k, v in inputs.items()}\n\n        with torch.no_grad():\n            outputs = model.generate(\n                **inputs,\n                max_new_tokens=50,\n                do_sample=False,\n                repetition_penalty=1.2,\n                no_repeat_ngram_size=3,\n            )\n\n        prediction = processor.decode(outputs[0], skip_special_tokens=True)\n        # Remove prefix from output\n        prediction = prediction.replace(\"<image>caption\", \"\").strip()\n        print(f\"🤖 Predicted: {prediction}\", flush=True)\n        print(f\"✅ Truth: {row['caption'][:80]}\", flush=True)\n\n    print(\"\\n\" + \"=\"*70, flush=True)\n    print(\"🎉 ALL DONE!\", flush=True)\n    print(\"=\"*70, flush=True)\n\nexcept Exception as e:\n    print(f\"\\n❌ TRAINING FAILED: {e}\", flush=True)\n    import traceback\n    traceback.print_exc()\n    raise\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}