TubeCLIP / predict.py
Krudev's picture
Upload predict.py
138d2a7 verified
Raw History Blame Contribute Delete
7.41 kB
# predict.py
import argparse
import os
from PIL import Image
import torch
import torch.nn as nn
import torch.nn.functional as F
import pytorch_lightning as pl
from transformers import CLIPProcessor, CLIPModel
from peft import LoraConfig, get_peft_model
from typing import Dict, Any
# =========================================================================
# Re-define the Model Class (must match your training script exactly)
# =========================================================================
# Global constants from your training script
NUM_CLASSES = 3
CLIP_MODEL_NAME = "openai/clip-vit-large-patch14"
LORA_R = 32
LORA_ALPHA = 64
LORA_DROPOUT = 0.1
LORA_TARGET_MODULES = ["q_proj", "k_proj", "v_proj", "out_proj", "fc1", "fc2"]
FEATURE_COMBINATION_STRATEGY = "mean"
LEARNING_RATE = 2e-4
CLASS_NAMES = ['10k-100k', '100k-1M', '1M+']
class CLIPForViewsClassification(pl.LightningModule):
def __init__(self, num_classes, clip_model_name, lora_r, lora_alpha, lora_dropout, lora_target_modules, learning_rate, feature_combination_strategy, num_training_steps_total):
super().__init__()
self.save_hyperparameters()
self.num_classes = num_classes
self.learning_rate = learning_rate
self.feature_combination_strategy = feature_combination_strategy
self.num_training_steps_total = num_training_steps_total
self.clip_model = CLIPModel.from_pretrained(clip_model_name)
lora_config = LoraConfig(r=lora_r, lora_alpha=lora_alpha, target_modules=lora_target_modules, lora_dropout=lora_dropout, bias="none")
self.clip_model.vision_model = get_peft_model(self.clip_model.vision_model, lora_config)
self.clip_model.text_model = get_peft_model(self.clip_model.text_model, lora_config)
embedding_dim = self.clip_model.config.projection_dim
classifier_input_dim = embedding_dim * 2 if self.feature_combination_strategy == "concat" else embedding_dim
self.classifier = nn.Sequential(
nn.Linear(classifier_input_dim, 512),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(512, self.num_classes)
)
self.processor = CLIPProcessor.from_pretrained(clip_model_name)
def forward(self, pixel_values, input_ids, attention_mask):
image_embeds = self.clip_model.get_image_features(pixel_values=pixel_values)
text_embeds = self.clip_model.get_text_features(input_ids=input_ids, attention_mask=attention_mask)
# --- THE CORRECTED FIX ---
# If the model returns a dataclass instead of a raw tensor,
# the final projected embeddings are ALREADY inside the pooler_output.
# We just extract them safely without projecting them a second time.
if not isinstance(image_embeds, torch.Tensor):
image_embeds = getattr(image_embeds, "pooler_output", image_embeds)
if not isinstance(text_embeds, torch.Tensor):
text_embeds = getattr(text_embeds, "pooler_output", text_embeds)
# -----------------------------------
# Now both are guaranteed to be standard PyTorch tensors
image_embeds = F.normalize(image_embeds, p=2, dim=-1)
text_embeds = F.normalize(text_embeds, p=2, dim=-1)
if self.feature_combination_strategy == "mean":
combined_features = (image_embeds + text_embeds) / 2
else:
combined_features = torch.cat((image_embeds, text_embeds), dim=-1)
logits = self.classifier(combined_features)
return logits
def predict_step(self, image_path: str, title: str) -> Dict[str, Any]:
image = Image.open(image_path).convert("RGB")
inputs = self.processor(text=[title], images=[image], return_tensors="pt", padding="max_length", truncation=True)
device = self.device
inputs = {k: v.to(device) for k, v in inputs.items()}
with torch.no_grad():
logits = self(inputs['pixel_values'], inputs['input_ids'], inputs['attention_mask'])
probabilities = F.softmax(logits, dim=1).squeeze(0)
predicted_class_idx = torch.argmax(probabilities).item()
predicted_class = CLASS_NAMES[predicted_class_idx]
return {
"image_path": image_path,
"title": title,
"predicted_class": predicted_class,
"probabilities": probabilities.cpu().numpy()
}
# =========================================================================
# Main Script Logic
# =========================================================================
def main():
parser = argparse.ArgumentParser(description="Use a fine-tuned CLIP model to predict YouTube view categories.")
parser.add_argument("--model_path", type=str, required=True,
help="Path to the trained model checkpoint (.ckpt).")
parser.add_argument("--input_path", type=str, required=True,
help="Path to a single image file or a directory of images to predict on.")
parser.add_argument("--title", type=str, default="A YouTube video",
help="The video title to be used for prediction. Default is 'A YouTube video'.")
args = parser.parse_args()
# Load the trained model
try:
model = CLIPForViewsClassification.load_from_checkpoint(
args.model_path,
num_classes=NUM_CLASSES,
clip_model_name=CLIP_MODEL_NAME,
lora_r=LORA_R,
lora_alpha=LORA_ALPHA,
lora_dropout=LORA_DROPOUT,
lora_target_modules=LORA_TARGET_MODULES,
learning_rate=LEARNING_RATE,
feature_combination_strategy=FEATURE_COMBINATION_STRATEGY,
num_training_steps_total=1000 # Placeholder, not used for inference
)
model.eval()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)
print(f"Model loaded successfully from {args.model_path} and moved to {device}.")
except Exception as e:
print(f"Error loading model: {e}")
return
# Check if input path is a file or directory
if os.path.isfile(args.input_path):
image_paths = [args.input_path]
elif os.path.isdir(args.input_path):
image_paths = [
os.path.join(args.input_path, f)
for f in os.listdir(args.input_path)
if f.lower().endswith(('.png', '.jpg', '.jpeg'))
]
if not image_paths:
print(f"No valid images found in directory: {args.input_path}")
return
else:
print(f"Error: Invalid input path. Must be a valid file or directory.")
return
# Make predictions
print("\n--- Predictions ---")
for img_path in image_paths:
try:
prediction_result = model.predict_step(img_path, args.title)
print(f"\nImage: {prediction_result['image_path']}")
print(f"Predicted Class: {prediction_result['predicted_class']}")
print(f"Probabilities: {prediction_result['probabilities']}")
except Exception as e:
print(f"Error processing image {img_path}: {e}")
if __name__ == "__main__":
main()