Download predict.py from Krudev/TubeCLIP: direct link, hf CLI and curl.
- Browser
- Download file 7.41 kB
-
https://huggingface.co/Krudev/TubeCLIP/resolve/main/predict.py
- Command line
-
hf download hf://Krudev/TubeCLIP/predict.py
-
curl -L -o predict.py https://huggingface.co/Krudev/TubeCLIP/resolve/main/predict.py
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() |