# 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()