# import
from transformers import AutoProcessor, AutoModel
import json
from PIL import Image
import torch
import os
import numpy as np
from tqdm import tqdm

# load model
device = "cuda"
processor_name_or_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
model_pretrained_name_or_path = "yuvalkirstain/PickScore_v1"

processor = AutoProcessor.from_pretrained(processor_name_or_path, cache_dir = '/cfs/cfs-1dafgugv/connorxian/hf_cache')
model = AutoModel.from_pretrained(model_pretrained_name_or_path, cache_dir = '/cfs/cfs-1dafgugv/connorxian/hf_cache').eval().to(device)

def calc_probs(prompt, images):
    
    # preprocess
    image_inputs = processor(
        images=images,
        padding=True,
        truncation=True,
        max_length=77,
        return_tensors="pt",
    ).to(device)
    
    text_inputs = processor(
        text=prompt,
        padding=True,
        truncation=True,
        max_length=77,
        return_tensors="pt",
    ).to(device)


    with torch.no_grad():
        # embed
        image_embs = model.get_image_features(**image_inputs)
        image_embs = image_embs / torch.norm(image_embs, dim=-1, keepdim=True)
    
        text_embs = model.get_text_features(**text_inputs)
        text_embs = text_embs / torch.norm(text_embs, dim=-1, keepdim=True)
    
        # score
        scores = model.logit_scale.exp() * (text_embs @ image_embs.T)[0]
        
        # get probabilities if you have multiple images to choose from
        probs = torch.softmax(scores, dim=-1)
    
    return scores.cpu().tolist(), probs.cpu().tolist()

data = json.load(open("/workspace/user_code/DiffusionDPO/flow_grpo/scripts/inference/result.json", "r"))
result = []
for i in tqdm(range(len(data))):
    item = data[i]
    pil_images = [Image.open(item['image']), Image.open(item['ori_image'])]
    # 如果pil_images是全黑的图片，则跳过
    if np.all(np.array(pil_images[0]) == 0) or np.all(np.array(pil_images[1]) == 0):
        # print(os.path.basename(item['image']))
        continue
    prompt = item["prompt"]
    score, prob = calc_probs(prompt, pil_images)
    result.append(score)

print(f'data number: {len(result)}')
result = np.array(result)
print(result.mean(axis=0))

