import torch
from transformers import CLIPProcessor, CLIPModel
from PIL import Image
import cv2, os
import pandas as pd
from tqdm import tqdm

# Load CLIP (smallest model for CPU)
model_id = "openai/clip-vit-base-patch32"
model = CLIPModel.from_pretrained(model_id)
processor = CLIPProcessor.from_pretrained(model_id)
model.eval()

labels = [
    "a person performing a barbell back squat",
    "a person performing a barbell front squat",
    "a person performing a bench press",
    "a person performing a deadlift",
    "a person performing an overhead press",
    "a person performing a bicep curl",
    "a person performing a pull-up",
    "a person performing a push-up",
    "a person performing a barbell row",
    "a person performing a lat pulldown",
    "a person performing a lunge",
    "a person performing a hip thrust",
    "a person performing a shoulder lateral raise"
]

video_folder = "videos"
results = []

def predict_video(video_path):
    cap = cv2.VideoCapture(video_path)
    frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
    frames = []

    # sample up to 12 evenly spaced frames
    for i in range(12):
        cap.set(cv2.CAP_PROP_POS_FRAMES, int(frame_count * i / 12))
        ret, frame = cap.read()
        if ret:
            frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
            frames.append(Image.fromarray(frame))
    cap.release()

    if not frames:
        return None, 0.0

    inputs = processor(text=labels, images=frames, return_tensors="pt", padding=True)
    with torch.no_grad():
        outputs = model(**inputs)
    logits_per_image = outputs.logits_per_image
    probs = logits_per_image.softmax(dim=1).mean(0)
    top = probs.argmax().item()
    return labels[top], probs[top].item()

for vid in tqdm(os.listdir(video_folder)):
    if vid.endswith((".mp4", ".mov", ".avi", ".mkv")):
        label, conf = predict_video(os.path.join(video_folder, vid))
        results.append({"filename": vid, "predicted_label": label, "confidence": conf})

pd.DataFrame(results).to_csv("auto_labels.csv", index=False)
print("Done! Saved predictions to auto_labels.csv")
