# -- Code Cell --
import os
y_json = [os.path.join("train", f"{i:05d}", "image.png") for i in range(200)]
y_json

# -- Code Cell --
import os
import json

all_data = []

for i in range(200):
    path = os.path.join("train", f"{i:05d}", "metadata.json")
    
    with open(path, "r", encoding="utf-8") as f:
        data = json.load(f)
        all_data.append(data)

# -- Code Cell --
all_data

# -- Code Cell --
from PIL import Image
images = []
for i in range(200):
    images.append(Image.open(f"./train/{i:05d}/image.png"))
images

# -- Code Cell --
import os
os.environ["HF_HOME"] = "E:/huggingface_cache"

# -- Code Cell --
from transformers import CLIPProcessor, CLIPModel
import torch
import torch.nn.functional as F
model = CLIPModel.from_pretrained("laion/CLIP-ViT-bigG-14-laion2B-39B-b160k")
processor = CLIPProcessor.from_pretrained("laion/CLIP-ViT-bigG-14-laion2B-39B-b160k")

# -- Code Cell --
device = torch.device("cuda")

# -- Code Cell --
import os, json
from PIL import Image

all_data_test = []
images_test = []

for i in range(len(os.listdir("./test"))):
    folder = os.path.join("test", f"{i:05d}")
    
    with open(os.path.join(folder, "metadata.json"), "r") as f:
        all_data_test.append(json.load(f))
    
    images_test.append(
        Image.open(os.path.join(folder, "image.png")).convert("RGB")
    )

# -- Code Cell --
templates = [
    "a photo of a {}",
    "an illustration of a {}",
    "a surreal drawing of a {}",
    "a magical depiction of a {}",
]

# -- Code Cell --
from tqdm import tqdm
import torch
import torch.nn.functional as F

model = model.to(device)
model.eval()

preds = []

for i in tqdm(range(len(all_data_test))):
    raw_words = all_data_test[i]["word_choices"]
    image = images_test[i]

    image_inputs = processor(images=image, return_tensors="pt")
    image_inputs = {k: v.to(device) for k, v in image_inputs.items()}

    with torch.no_grad():
        image_outputs = model.vision_model(
            pixel_values=image_inputs["pixel_values"]
        )
        image_features = model.visual_projection(image_outputs.pooler_output)
        image_features = F.normalize(image_features, dim=-1)

    all_scores = torch.zeros(len(raw_words), device=device)

    for template in templates:
        words = [template.format(w) for w in raw_words]

        text_inputs = processor(
            text=words,
            return_tensors="pt",
            padding=True,
            truncation=True
        )
        text_inputs = {k: v.to(device) for k, v in text_inputs.items()}

        with torch.no_grad():
            text_outputs = model.text_model(
                input_ids=text_inputs["input_ids"],
                attention_mask=text_inputs["attention_mask"]
            )
            text_features = model.text_projection(text_outputs.pooler_output)
            text_features = F.normalize(text_features, dim=-1)

            similarity = (image_features @ text_features.T).squeeze(0)
            all_scores += similarity

    idx = all_scores.topk(k=5).indices.cpu().tolist()
    top5 = [raw_words[j] for j in idx]
    preds.append(top5)

# -- Code Cell --
import pandas as pd

df = pd.DataFrame(preds, columns=["word1","word2","word3","word4","word5"])
df.insert(0, "ID", [f"{i:05d}" for i in range(len(preds))])
df.to_csv("subi.csv", index=False)

# -- Code Cell --
