{"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":29705,"databundleVersionId":2662435,"sourceType":"competition"},{"sourceId":2793532,"sourceType":"datasetVersion","datasetId":1705878},{"sourceId":14197126,"sourceType":"datasetVersion","datasetId":9053625}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install git+https://github.com/openai/CLIP.git\n!pip install transformers huggingface-hub ftfy regex tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:16:43.093812Z","iopub.execute_input":"2025-12-20T16:16:43.094186Z","iopub.status.idle":"2025-12-20T16:16:56.730753Z","shell.execute_reply.started":"2025-12-20T16:16:43.094149Z","shell.execute_reply":"2025-12-20T16:16:56.729844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport base64\nimport io\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom transformers import AutoTokenizer, AutoModel, AutoConfig\nimport clip\nfrom torch.nn import functional as F\nfrom torchvision import transforms\n# 設定裝置\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:16:56.732450Z","iopub.execute_input":"2025-12-20T16:16:56.732781Z","iopub.status.idle":"2025-12-20T16:17:15.404435Z","shell.execute_reply.started":"2025-12-20T16:16:56.732748Z","shell.execute_reply":"2025-12-20T16:17:15.403769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DS_PATH = \"/kaggle/input/wikipedia-train-0\"  # 請確認這路徑是對的\nfilenames = sorted(os.listdir(DS_PATH))\njson_content = []\n\nprint(f\"🚀 正在讀取 JSON 資料集: {DS_PATH}...\")\nfor file in tqdm(filenames, desc=\"Loading Files\"):\n    if not file.endswith('.json'): continue\n    filename = os.path.join(DS_PATH, file)\n    with open(filename, \"rb\") as fr:\n        for line in fr:\n            if line:\n                obj = json.loads(line)\n                # 簡單檢查欄位\n                if \"b64_bytes\" in obj and \"wit_features\" in obj:\n                    # 提取所有可用的描述\n                    descriptions = []\n                    for element in obj[\"wit_features\"]:\n                        desc = element.get(\"caption_title_and_reference_description\")\n                        if desc:\n                            descriptions.append(desc)\n                    \n                    if descriptions and obj[\"b64_bytes\"]:\n                        # 為了簡化，我們這裡只拿第一個描述當作正樣本\n                        # (進階版可以把所有描述都拆出來當多筆資料)\n                        json_content.append({\n                            \"b64_bytes\": obj[\"b64_bytes\"],\n                            \"caption\": descriptions[0],\n                            \"url\": obj.get(\"image_url\", \"\")\n                        })\n\nprint(f\"✅ 資料讀取完成！共有 {len(json_content)} 筆圖文資料。\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:17:15.405331Z","iopub.execute_input":"2025-12-20T16:17:15.405747Z","iopub.status.idle":"2025-12-20T16:19:05.948999Z","shell.execute_reply.started":"2025-12-20T16:17:15.405696Z","shell.execute_reply":"2025-12-20T16:19:05.948099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================================\n# 1. 定義輔助類別 (零件)\n# ==========================================\nclass ContrastiveLoss(nn.Module):\n    def __init__(self, margin=0.2, max_violation=True): super().__init__()\n    def forward(self, x, y): return 0\n\nclass TextExtractorModel(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        text_model = config['text-model']['model-name']\n        self.finetune = config['text-model']['finetune']\n        self.text_model = AutoModel.from_pretrained(text_model)\n    def forward(self, ids, mask):\n        with torch.set_grad_enabled(self.finetune):\n            out = self.text_model(input_ids=ids, attention_mask=mask, output_hidden_states=True)\n        out = torch.stack(out.hidden_states, dim=0)\n        return out\n\nclass ImageExtractorModel(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.finetune = config['image-model']['finetune']\n        model_name = config['image-model']['model-name']\n        # 強制轉 FP32 避免 bug\n        self.clip_model, _ = clip.load(model_name, device='cpu')\n        self.clip_model = self.clip_model.float()\n    def forward(self, img):\n        with torch.set_grad_enabled(self.finetune):\n            feats = self.clip_model.encode_image(img)\n        return feats\n\nclass TransformerPooling(nn.Module):\n    def __init__(self, input_dim=1024, output_dim=1024, num_layers=2):\n        super().__init__()\n        transformer_layer = nn.TransformerEncoderLayer(d_model=input_dim, nhead=4, dim_feedforward=input_dim, dropout=0.1, activation='relu')\n        self.transformer_encoder = nn.TransformerEncoder(transformer_layer, num_layers=num_layers)\n        self.proj = nn.Linear(input_dim, output_dim) if input_dim != output_dim else None\n    def forward(self, input, mask):\n        mask_bool = ~mask.bool()\n        input = input.permute(1, 0, 2)\n        output = self.transformer_encoder(input, src_key_padding_mask=mask_bool)\n        output = output[0]\n        if self.proj: output = self.proj(output)\n        return output\n\nclass DepthAggregatorModel(nn.Module):\n    def __init__(self, aggr, input_dim=1024, output_dim=1024):\n        super().__init__()\n        self.aggr = aggr\n        if self.aggr == 'gated':\n            self.self_attn = nn.MultiheadAttention(input_dim, num_heads=4, dropout=0.1)\n            self.gate_ffn = nn.Linear(input_dim, 1)\n        self.proj = nn.Linear(input_dim, output_dim) if input_dim != output_dim else None\n    def forward(self, x, mask):\n        if self.aggr is None: out = x[-1, :, 0, :]\n        elif self.aggr == 'mean': out = x[:, :, 0, :].mean(dim=0)\n        if self.proj: out = self.proj(out)\n        return out\n\nclass FeatureFusionModel(nn.Module):\n    def __init__(self, mode, img_feat_dim, txt_feat_dim, common_space_dim):\n        super().__init__()\n        self.mode = mode\n        if mode == 'weighted':\n            self.alphas = nn.Sequential(\n                nn.Linear(img_feat_dim + txt_feat_dim, 512), nn.ReLU(), nn.Dropout(p=0.1), nn.Linear(512, 2))\n            self.img_proj = nn.Linear(img_feat_dim, common_space_dim)\n            self.txt_proj = nn.Linear(txt_feat_dim, common_space_dim)\n            self.post_process = nn.Sequential(\n                nn.Linear(common_space_dim, common_space_dim), nn.ReLU(), nn.Dropout(p=0.1), nn.Linear(common_space_dim, common_space_dim)\n            )\n    def forward(self, img_feat, txt_feat):\n        concat_feat = torch.cat([img_feat, txt_feat], dim=1)\n        alphas = torch.sigmoid(self.alphas(concat_feat))\n        img_feat_norm = F.normalize(self.img_proj(img_feat), p=2, dim=1)\n        txt_feat_norm = F.normalize(self.txt_proj(txt_feat), p=2, dim=1)\n        out_feat = img_feat_norm * alphas[:, 0].unsqueeze(1) + txt_feat_norm * alphas[:, 1].unsqueeze(1)\n        out_feat = self.post_process(out_feat)\n        return out_feat, alphas\n\n# ==========================================\n# 2. 核心模型 MatchingModel (已改裝 Inference 接口)\n# ==========================================\nclass MatchingModel(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        common_space_dim = config['matching']['common-space-dim']\n        num_text_transformer_layers = config['matching']['text-transformer-layers']\n        img_feat_dim = config['image-model']['dim']\n        txt_feat_dim = config['text-model']['dim']\n        image_disabled = config['image-model'].get('disabled', False)\n        \n        self.aggregate_tokens_depth = config['matching'].get('aggregate-tokens-depth', None)\n        self.fusion_mode = config['matching'].get('fusion-mode', 'concat')\n        self.image_disabled = image_disabled\n\n        self.txt_model = TextExtractorModel(config)\n        \n        if not image_disabled:\n            self.img_model = ImageExtractorModel(config)\n            self.image_fc = nn.Sequential(\n                nn.Linear(img_feat_dim, img_feat_dim), nn.Dropout(0.2), nn.ReLU(), nn.Linear(img_feat_dim, img_feat_dim)\n            )\n            if self.fusion_mode == 'concat':\n                self.process_after_concat = nn.Sequential(\n                    nn.Linear(img_feat_dim + txt_feat_dim, common_space_dim),\n                    nn.ReLU(), nn.Dropout(0.1),\n                    nn.Linear(common_space_dim, common_space_dim)\n                )\n            else:\n                self.process_after_concat = FeatureFusionModel(self.fusion_mode, img_feat_dim, txt_feat_dim, common_space_dim)\n\n        self.caption_process = TransformerPooling(txt_feat_dim, common_space_dim, num_text_transformer_layers)\n        self.url_process = TransformerPooling(txt_feat_dim, txt_feat_dim if not image_disabled else common_space_dim, num_text_transformer_layers)\n        if self.aggregate_tokens_depth:\n            self.token_aggregator = DepthAggregatorModel(self.aggregate_tokens_depth, txt_feat_dim, common_space_dim)\n        self.matching_loss = ContrastiveLoss()\n\n    # --- 這是 Inference 必備的接口 (我幫你從 compute_embeddings 拆出來的) ---\n    def encode_query(self, img, url, url_mask):\n        # 1. 計算 URL 文字特徵 (Test 時是 Dummy)\n        url_feats = self.txt_model(url, url_mask)\n        url_feats_plus = self.url_process(url_feats[-1], url_mask)\n        if self.aggregate_tokens_depth:\n            url_feats = url_feats_plus + self.token_aggregator(url_feats, url_mask)\n        else:\n            url_feats = url_feats_plus\n\n        # 2. 計算圖片特徵並融合\n        if not self.image_disabled:\n            img_feats = self.image_fc(self.img_model(img).float())\n            if self.fusion_mode == 'concat':\n                query_feats = torch.cat([img_feats, url_feats], dim=1)\n                query_feats = self.process_after_concat(query_feats)\n            else:\n                query_feats, _ = self.process_after_concat(img_feats, url_feats)\n        else:\n            query_feats = url_feats\n        \n        return F.normalize(query_feats, p=2, dim=1)\n\n    def encode_caption(self, caption, caption_mask):\n        caption_feats = self.txt_model(caption, caption_mask)\n        caption_feats_plus = self.caption_process(caption_feats[-1], caption_mask)\n        if self.aggregate_tokens_depth:\n            caption_feats = caption_feats_plus + self.token_aggregator(caption_feats, caption_mask)\n        else:\n            caption_feats = caption_feats_plus\n        return F.normalize(caption_feats, p=2, dim=1)\n\n# ==========================================\n# 3. 對應的 Config (必須跟 A 的設定一樣)\n# ==========================================\nConfig = {\n    'text-model': {'model-name': 'xlm-roberta-base', 'dim': 768, 'finetune': False},\n    'image-model': {'model-name': 'ViT-B/32', 'dim': 512, 'finetune': False, 'disabled': False},\n    'matching': {\n        'common-space-dim': 768, \n        'text-transformer-layers': 2, \n        'fusion-mode': 'concat', \n        'aggregate-tokens-depth': 'mean' # 這是關鍵！舊的 Config 可能沒有這個\n    },\n    'training': {'margin': 0.2, 'max-violation': False}\n}\nconfig = Config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:19:05.950638Z","iopub.execute_input":"2025-12-20T16:19:05.951243Z","iopub.status.idle":"2025-12-20T16:19:05.972641Z","shell.execute_reply.started":"2025-12-20T16:19:05.951218Z","shell.execute_reply":"2025-12-20T16:19:05.972103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 定義標準的 CLIP Normalization 數值 (這是固定的，不能改)\nnormalize = transforms.Normalize(mean=(0.48145466, 0.4578275, 0.40821073), \n                                 std=(0.26862954, 0.26130258, 0.27577711))\n\nTOKENIZER = AutoTokenizer.from_pretrained('xlm-roberta-base')\nMAX_LEN = 128\n\nclass JsonMiningDataset(Dataset):\n    def __init__(self, data_list, tokenizer, max_len):\n        self.data = data_list\n        self.tokenizer = tokenizer\n        self.max_len = max_len\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        item = self.data[idx]\n        \n        # 1. 圖片處理 (Base64 解碼 + 標準化)\n        try:\n            decoded = base64.b64decode(item['b64_bytes'])\n            image = Image.open(io.BytesIO(decoded)).convert(\"RGB\")\n            image = image.resize((224, 224))\n            \n            # 轉 Tensor 並除以 255\n            image = torch.tensor(np.array(image)).permute(2, 0, 1).float() / 255.0\n            \n            # ⚠️【關鍵修改】加上這行標準化！\n            image = normalize(image) \n        except:\n            image = torch.zeros(3, 224, 224)\n            \n        # 2. 文字處理\n        caption = str(item['caption'])\n        inputs = self.tokenizer.encode_plus(\n            caption, None, add_special_tokens=True,\n            max_length=self.max_len, padding='max_length', truncation=True, return_tensors='pt'\n        )\n        \n        return {\n            'image': image,\n            'input_ids': inputs['input_ids'].flatten(),\n            'attention_mask': inputs['attention_mask'].flatten(),\n            'caption_text': caption, \n            'index': idx \n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:19:05.973489Z","iopub.execute_input":"2025-12-20T16:19:05.973697Z","iopub.status.idle":"2025-12-20T16:19:08.429267Z","shell.execute_reply.started":"2025-12-20T16:19:05.973678Z","shell.execute_reply":"2025-12-20T16:19:08.428665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# A. 初始化模型\n# (注意：這裡的 config 來自 Cell 4 的定義，請確保 Cell 4 已經執行過)\ncoarse_model = MatchingModel(config)\n\n# B. 載入 A 的權重\n# ⚠️ 請確認這個路徑是你上傳的那個 .bin 檔案\nWEIGHTS_PATH = \"/kaggle/input/mining/Loss_2.5559_epoch_9.bin\" \n\ntry:\n    # 這裡絕對不能加 strict=False，我們要確保它真的完美載入！\n    coarse_model.load_state_dict(torch.load(WEIGHTS_PATH, map_location=device))\n    print(\"✅ 恭喜！模型架構終於對上了！權重完美載入！\")\n    coarse_model.to(device).eval()\nexcept Exception as e:\n    print(\"❌ 還是有錯... 請把下面的錯誤訊息給我：\")\n    print(e)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:19:08.430154Z","iopub.execute_input":"2025-12-20T16:19:08.430463Z","iopub.status.idle":"2025-12-20T16:20:01.334300Z","shell.execute_reply.started":"2025-12-20T16:19:08.430440Z","shell.execute_reply":"2025-12-20T16:20:01.333131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# A. 載入模型\n\nmodel = MatchingModel(config)\n\nWEIGHTS_PATH = \"/kaggle/input/mining/Loss_2.5559_epoch_9.bin\" # 請改路徑\n\nmodel.load_state_dict(torch.load(WEIGHTS_PATH, map_location=device))\n\nmodel.to(device)\n\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:20:01.335641Z","iopub.execute_input":"2025-12-20T16:20:01.336425Z","iopub.status.idle":"2025-12-20T16:20:08.423110Z","shell.execute_reply.started":"2025-12-20T16:20:01.336386Z","shell.execute_reply":"2025-12-20T16:20:08.422508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# B. 建立 DataLoader\n\nmining_dataset = JsonMiningDataset(json_content, TOKENIZER, MAX_LEN) \nmining_loader = DataLoader(mining_dataset, batch_size=64, shuffle=False)\n\nall_img_embs = []\nall_txt_embs = []\nall_captions = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:20:08.424104Z","iopub.execute_input":"2025-12-20T16:20:08.424591Z","iopub.status.idle":"2025-12-20T16:20:08.428434Z","shell.execute_reply.started":"2025-12-20T16:20:08.424568Z","shell.execute_reply":"2025-12-20T16:20:08.427692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# C. 計算向量\nwith torch.no_grad():\n    for data in tqdm(mining_loader, desc=\"Encoding\"):\n        images = data['image'].to(device)\n        ids = data['input_ids'].to(device)\n        mask = data['attention_mask'].to(device)\n        \n        q_feat = model.encode_query(images, ids, mask)\n        c_feat = model.encode_caption(ids, mask)\n        \n        all_img_embs.append(q_feat.cpu())\n        all_txt_embs.append(c_feat.cpu())\n        all_captions.extend(data['caption_text']) # 把文字存起來\n\nall_img_embs = torch.cat(all_img_embs, dim=0).to(device)\nall_txt_embs = torch.cat(all_txt_embs, dim=0).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T16:20:08.429331Z","iopub.execute_input":"2025-12-20T16:20:08.429604Z","execution_failed":"2025-12-20T16:21:07.213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# D. 找錯題\nhard_negatives_data = []\nbatch_size = 1000\n\nfor i in tqdm(range(0, len(all_img_embs), batch_size), desc=\"Mining\"):\n    end = min(i + batch_size, len(all_img_embs))\n    batch_img = all_img_embs[i:end]\n    \n    # 算相似度\n    sim_matrix = torch.matmul(batch_img, all_txt_embs.T)\n    \n    # 取 Top-10\n    vals, indices = torch.topk(sim_matrix, k=10, dim=1)\n    indices = indices.cpu().numpy()\n    \n    for idx_in_batch, candidates in enumerate(indices):\n        real_idx = i + idx_in_batch\n        positive_caption = all_captions[real_idx]\n        \n        neg_list = []\n        for cand_id in candidates:\n            if cand_id != real_idx: # 不是正確答案\n                neg_caption = all_captions[cand_id]\n                \n                # 存起來\n                hard_negatives_data.append({\n                    'image_id': real_idx, # 這裡存 index 方便對應，進階可以存 base64\n                    'positive': positive_caption,\n                    'negative': neg_caption,\n                    'rank': len(neg_list) + 1\n                })\n                neg_list.append(cand_id)\n            if len(neg_list) >= 2: break # 每張圖挖 2 個錯題","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# E. 存檔\nmining_df = pd.DataFrame(hard_negatives_data)\nmining_df.to_csv(\"train_hard_negatives.csv\", index=False)\nprint(\"挖礦完成！已產出 train_hard_negatives.csv\")\nprint(mining_df.head())","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\n# 刪除不用的變數\ndel model, mining_loader, mining_dataset, all_img_embs, all_txt_embs\ntorch.cuda.empty_cache()\ngc.collect()\nprint(\"♻️ 記憶體已釋放，準備開始訓練精排模型...\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom transformers import XLMRobertaForSequenceClassification\nfrom torch.optim import AdamW \nimport os\n# 1. 定義 Dataset\nclass RerankDataset(Dataset):\n    def __init__(self, mining_df, tokenizer, max_len):\n        self.df = mining_df\n        self.tokenizer = tokenizer\n        self.max_len = max_len\n        \n    def __len__(self):\n        return len(self.df) * 2 # 正樣本 + 負樣本\n    \n    def __getitem__(self, idx):\n        # 偶數 index 是正樣本 (label=1)\n        # 奇數 index 是負樣本 (label=0)\n        row_idx = idx // 2\n        row = self.df.iloc[row_idx]\n        \n        if idx % 2 == 0:\n            text_a = row['positive'] \n            text_b = row['positive'] \n            label = 1.0\n        else:\n            text_a = row['positive']\n            text_b = row['negative']\n            label = 0.0\n            \n        # Cross-Encoder 的輸入是把兩句話接在一起\n        inputs = self.tokenizer.encode_plus(\n            text_a, text_b, # 兩句話\n            add_special_tokens=True,\n            max_length=self.max_len, padding='max_length', truncation=True, return_tensors='pt'\n        )\n        \n        return {\n            'input_ids': inputs['input_ids'].flatten(),\n            'attention_mask': inputs['attention_mask'].flatten(),\n            'labels': torch.tensor(label, dtype=torch.float)\n        }","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. 準備資料\nmining_df = pd.read_csv(\"train_hard_negatives.csv\")\n\n\nrerank_dataset = RerankDataset(mining_df, TOKENIZER, max_len=128) \nrerank_loader = DataLoader(rerank_dataset, batch_size=16, shuffle=True)\n\n# 3. 定義模型 (使用 XLM-R 做二元分類)\nprint(\"🚀 初始化 Cross-Encoder 模型...\")\nrerank_model = XLMRobertaForSequenceClassification.from_pretrained('xlm-roberta-base', num_labels=1)\nrerank_model.to(device)\nrerank_model.train()\n\n# 記得 AdamW 要從 torch.optim 匯入 (如上一則回答所述)\nfrom torch.optim import AdamW\noptimizer = AdamW(rerank_model.parameters(), lr=2e-5)\ncriterion = nn.BCEWithLogitsLoss()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nepochs = 1\n\nprint(f\"🔥 開始訓練精排模型 (Epochs={epochs}, 使用 AMP 加速)...\")\n\n# 初始化 Scaler (加速神器)\nscaler = GradScaler()\n\nfor epoch in range(epochs):\n    total_loss = 0\n    rerank_model.train() # 確保在訓練模式\n    \n    # 顯示進度條\n    loop = tqdm(rerank_loader, desc=f\"Epoch {epoch+1}/{epochs}\")\n    \n    for batch in loop:\n        ids = batch['input_ids'].to(device)\n        mask = batch['attention_mask'].to(device)\n        labels = batch['labels'].to(device)\n        \n        optimizer.zero_grad()\n        \n        # ⚡ 開啟混合精度計算 (Forward)\n        with autocast():\n            outputs = rerank_model(ids, attention_mask=mask)\n            logits = outputs.logits.squeeze()\n            loss = criterion(logits, labels)\n        \n        # ⚡ 使用 Scaler 反向傳播 (Backward)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        \n        # 更新進度條上的 Loss\n        loop.set_postfix(loss=f\"{loss.item():.4f}\")\n        \n    avg_loss = total_loss / len(rerank_loader)\n    print(f\"Epoch {epoch+1} Average Loss: {avg_loss:.4f}\")\n\n\ntorch.save(rerank_model.state_dict(), \"reranker_model.bin\")\nprint(\"模型已存檔為 reranker_model.bin\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import AutoTokenizer, XLMRobertaForSequenceClassification\nimport glob\nimport base64\nimport io\nfrom PIL import Image\nfrom tqdm.auto import tqdm\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nTOKENIZER = AutoTokenizer.from_pretrained('xlm-roberta-base')\nMAX_LEN = 64\n\n# 路徑設定\nTEST_TSV = '/kaggle/input/c/wikipedia-image-caption/test.tsv'\nCAPTION_CSV = '/kaggle/input/c/wikipedia-image-caption/test_caption_list.csv'\nPIXEL_DIR = '/kaggle/input/c/wikipedia-image-caption/image_data_test/image_pixels'\n\n# 粗排模型權重 (用來撈候選人)\nCOARSE_WEIGHTS = \"/kaggle/input/mining/Loss_2.5559_epoch_9.bin\" \n# 精排模型權重 (剛訓練好的)\nRERANK_WEIGHTS = \"reranker_model.bin\" \n\nprint(f\"Device: {DEVICE}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. 定義兩個模型 (粗排 & 精排)\n# ==========================================\n# --- 粗排模型  ---\n# (請確保前面已經定義過 MatchingModel, TextExtractorModel 等類別)\n\ncoarse_model = MatchingModel(config) \ncoarse_model.load_state_dict(torch.load(COARSE_WEIGHTS, map_location=DEVICE)\n                            )\ncoarse_model.to(DEVICE).eval()\nprint(\"粗排模型載入完成\")\n\n# --- 精排模型 (你的架構) ---\nrerank_model = XLMRobertaForSequenceClassification.from_pretrained('xlm-roberta-base', num_labels=1)\nrerank_model.load_state_dict(torch.load(RERANK_WEIGHTS, map_location=DEVICE))\nrerank_model.to(DEVICE).eval()\nprint(\"精排模型載入完成\")\n\n# ==========================================\n# 3. 準備測試資料\n# ==========================================\n# 讀取測試圖片 (Base64)\nprint(\"正在載入測試圖片庫...\")\nimage_map = {}\npixel_files = sorted(glob.glob(f\"{PIXEL_DIR}/*.csv\"))\nfor f in tqdm(pixel_files):\n    temp = pd.read_csv(f, sep='\\t', names=['url', 'b64'], usecols=[0,1])\n    for _, row in temp.iterrows():\n        image_map[row['url']] = row['b64']\n\n# 讀取候選標題\ncaptions_df = pd.read_csv(CAPTION_CSV)\nall_captions = captions_df['caption_title_and_reference_description'].tolist()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"normalize = transforms.Normalize(mean=(0.48145466, 0.4578275, 0.40821073), \n                                 std=(0.26862954, 0.26130258, 0.27577711))\n\nclass TestImageDataset(Dataset):\n    def __init__(self, tsv_path, img_map, tokenizer):\n        self.df = pd.read_csv(tsv_path, sep='\\t')\n        self.img_map = img_map\n        self.tokenizer = tokenizer\n        \n    def __len__(self): return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        url = row['image_url']\n        b64 = self.img_map.get(url, \"\")\n        \n        # 1. 圖片處理\n        try:\n            image = Image.open(io.BytesIO(base64.b64decode(b64))).convert(\"RGB\")\n            image = image.resize((224, 224))\n            image = torch.tensor(np.array(image)).permute(2, 0, 1).float() / 255.0\n            image = normalize(image)\n        except:\n            image = torch.zeros(3, 224, 224)\n        \n        # 2. 關鍵修正：從 URL 提取檔名作為文字輸入\n        # 原本是 dummy_input = \"\" \n        try:\n            # 取網址最後一段 -> 解碼 (%20變空白) -> 去掉副檔名 -> 底線變空白\n            filename = url.split('/')[-1] \n            filename = urllib.parse.unquote(filename)\n            filename = filename.replace('_', ' ').replace('.jpg', '').replace('.png', '')\n            text_input = filename\n        except:\n            text_input = \"\"\n            \n        # 3. 編碼文字\n        tokenized_input = self.tokenizer.encode_plus(\n            text_input, \n            max_length=64, \n            padding='max_length', \n            truncation=True, \n            return_tensors='pt'\n        )\n        \n        return {\n            'image': image, \n            'input_ids': tokenized_input['input_ids'].flatten(),\n            'attention_mask': tokenized_input['attention_mask'].flatten(),\n            'id': row['id']\n        }\n\nclass CapDataset(Dataset):\n    def __init__(self, caps, tokenizer):\n        self.caps = caps\n        self.tokenizer = tokenizer\n    def __len__(self): return len(self.caps)\n    def __getitem__(self, idx):\n        inputs = self.tokenizer.encode_plus(str(self.caps[idx]), return_tensors='pt', max_length=64, padding='max_length', truncation=True)\n        return inputs['input_ids'].flatten(), inputs['attention_mask'].flatten()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. 計算向量 (Encoding)\n# ==========================================\nprint(\"[Step 1] 準備測試資料與計算向量...\")\n\n# --- 2.1 計算圖片向量 ---\ntest_ds = TestImageDataset(TEST_TSV, image_map, TOKENIZER)\ntest_loader = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=0)\n\nimg_embs = []\ntest_ids = []\n\nwith torch.no_grad():\n    for data in tqdm(test_loader, desc=\"Encoding Images\"):\n        img = data['image'].to(DEVICE)\n        ids = data['input_ids'].to(DEVICE)\n        mask = data['attention_mask'].to(DEVICE)\n        \n        feat = coarse_model.encode_query(img, ids, mask)\n        img_embs.append(feat.cpu())\n        test_ids.extend(data['id'].numpy())\n\nimg_embs = torch.cat(img_embs, dim=0)\nprint(f\"圖片向量計算完成 形狀: {img_embs.shape}\")\n\n# --- 2.2 計算文字向量 ---\n\ncap_loader = DataLoader(CapDataset(all_captions, TOKENIZER), batch_size=256, shuffle=False, num_workers=0)\n\ntxt_embs = []\nwith torch.no_grad():\n    for ids, mask in tqdm(cap_loader, desc=\"Encoding Captions\"):\n        ids, mask = ids.to(DEVICE), mask.to(DEVICE)\n        feat = coarse_model.encode_caption(ids, mask)\n        txt_embs.append(feat.cpu())\n\ntxt_embs = torch.cat(txt_embs, dim=0)\nprint(f\"文字向量計算完成 形狀: {txt_embs.shape}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================================\n# 3. 粗排海選 (Coarse Retrieval) - 算出 Top-50\n# ==========================================\nprint(\"[Step 2] 粗排...\")\n\n# 確保在 CPU (避免 GPU OOM)\nif img_embs.is_cuda: img_embs = img_embs.cpu()\nif txt_embs.is_cuda: txt_embs = txt_embs.cpu()\n\n# 文字向量搬到 GPU\ntxt_embs_gpu = txt_embs.to(DEVICE)\n\ntop100_indices_list = []\nBATCH_SIZE = 1000 \n\nfor i in tqdm(range(0, len(img_embs), BATCH_SIZE), desc=\"Coarse Retrieval\"):\n    batch_img = img_embs[i : i + BATCH_SIZE].to(DEVICE)\n    batch_sims = torch.matmul(batch_img, txt_embs_gpu.T)\n    _, batch_topk = torch.topk(batch_sims, k=100 , dim=1)\n    top100_indices_list.append(batch_topk.cpu())\n    del batch_img, batch_sims, batch_topk\n\ntop100_indices = torch.cat(top100_indices_list, dim=0)\nprint(f\"粗排完成{top100_indices.shape}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 隨機抽查一張測試圖片，看看是不是黑畫面\nimport matplotlib.pyplot as plt\n\n# 拿 Test Dataset 的第 0 筆資料\ndataset_check = TestImageDataset(TEST_TSV, image_map, TOKENIZER)\ndata = dataset_check[0]\n\nimg_tensor = data['image']\nprint(f\"圖片 Tensor 形狀: {img_tensor.shape}\")\nprint(f\"圖片數值範圍: Min={img_tensor.min():.4f}, Max={img_tensor.max():.4f}, Mean={img_tensor.mean():.4f}\")\n\n# 如果 Min=0, Max=0，代表你讀到黑畫面了 (圖片解碼失敗)\nif img_tensor.max() == 0:\n    print(\"❌ 警告：這張圖是全黑的！你的 image_map 或 base64 解碼有問題！\")\nelse:\n    print(\"✅ 圖片數值正常 (不是全黑)。\")\n\n# 嘗試把 Tensor 轉回圖片顯示 (因為有 Normalize，顏色會怪怪的，但要有東西)\n# 反標準化\ninv_normalize = transforms.Normalize(\n    mean=[-0.48145466/0.26862954, -0.4578275/0.26130258, -0.40821073/0.27577711],\n    std=[1/0.26862954, 1/0.26130258, 1/0.27577711]\n)\nimg_display = inv_normalize(img_tensor).permute(1, 2, 0).numpy()\nimg_display = np.clip(img_display, 0, 1)\n\nplt.imshow(img_display)\nplt.title(f\"Check ID: {data['id']}\")\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 隨機挑 5 張圖，印出它們粗排第一名的標題\nprint(\"🔍 檢查粗排品質 (Top-1 預測):\")\nfor i in range(5):\n    idx = top100_indices[i, 0] # 取第 i 張圖的第 1 名候選人索引\n    print(f\"圖 {i} 的粗排第一名: {all_captions[idx]}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================================\n# 4. 精排決選 (高速加速版 + Top-100 適配)\n# ==========================================\nfrom torch.cuda.amp import autocast # 匯入加速工具\nfrom torch.utils.data import TensorDataset\n\nprint(\"[Step 3] 精排：開始重新打分...\")\n\nfinal_submission = []\nALPHA = 0.5\n\n# 增加 Batch Size 到 256 (根據顯存情況調整，T4 跑純文字通常 256 沒問題)\n# 如果報錯 OOM，請改回 128 或 64\nRERANK_BATCH_SIZE = 256 \n\nfor i, candidates in enumerate(tqdm(top100_indices.cpu().numpy(), desc=\"Reranking\")):\n    img_id = test_ids[i]\n    \n    # 1. 準備 Query 和 Candidates\n    pseudo_query = all_captions[candidates[0]] \n    candidate_texts = [all_captions[idx] for idx in candidates]\n    \n    # 2. 準備 Batch 資料\n    pairs = [[pseudo_query, cand] for cand in candidate_texts]\n    encoded = TOKENIZER(pairs, padding=True, truncation=True, max_length=64, return_tensors=\"pt\")\n    \n    # 建立一個臨時的 DataLoader 來做 Batch 推論 (避免一次塞 100 個爆掉，雖然 100 個通常還好)\n    # 但為了穩健，我們還是乖乖切 Batch\n    batch_input_ids = encoded['input_ids']\n    batch_attention_mask = encoded['attention_mask']\n    \n    # 建立小型的 Dataset\n    mini_dataset = TensorDataset(batch_input_ids, batch_attention_mask)\n    mini_loader = DataLoader(mini_dataset, batch_size=RERANK_BATCH_SIZE, shuffle=False)\n    \n    all_scores = []\n    \n    # 3. 執行推論\n    with torch.no_grad():\n        with autocast(): # ⚡ 開啟混合精度加速 (關鍵!)\n            for b_ids, b_mask in mini_loader:\n                b_ids = b_ids.to(DEVICE)\n                b_mask = b_mask.to(DEVICE)\n                \n                outputs = rerank_model(input_ids=b_ids, attention_mask=b_mask)\n                scores = torch.sigmoid(outputs.logits.squeeze()).cpu().numpy()\n                \n                # 處理 scores 可能是純量(Scalar)的情況 (當 batch=1)\n                if scores.ndim == 0: scores = [scores]\n                all_scores.extend(scores)\n    \n    rerank_scores = np.array(all_scores)\n    \n    # 4. 取得粗排分數\n    coarse_scores = np.linspace(1.0, 0.0, len(candidates))\n    \n    # 5. 分數融合\n    final_scores = (ALPHA * coarse_scores) + ((1 - ALPHA) * rerank_scores)\n    \n    # 6. 重新排序\n    reranked_order = np.argsort(final_scores)[::-1]\n    \n    # 取前 5 名\n    top5_local_indices = reranked_order[:5]\n    top5_global_indices = [candidates[k] for k in top5_local_indices]\n    \n    for rank, idx in enumerate(top5_global_indices):\n        final_submission.append({\n            \"id\": img_id,\n            \"caption_title_and_reference_description\": all_captions[idx]\n        })\n\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================================\n# 5. 存檔\n# ==========================================\npd.DataFrame(final_submission).to_csv(\"submission.csv\", index=False)\nprint(\"🎉 submission.csv 已生成\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-20T16:21:07.215Z"}},"outputs":[],"execution_count":null}]}