Views
No views yet
1# Sample code for loading and using this model
2from transformers import AutoProcessor, AutoModelForCausalLM
3from peft import PeftModel
4import torch
5from PIL import Image
6
7# Load base model and processor
8base_model_id = "unsloth/llama-3.2-11b-vision-instruct"
9processor = AutoProcessor.from_pretrained(base_model_id)
10model = AutoModelForCausalLM.from_pretrained(base_model_id, device_map="auto")
11
12# Load this adapter
13adapter_id = "saakshigupta/deepfake-explainer-2"
14model = PeftModel.from_pretrained(model, adapter_id)
15
16# Function to fix cross-attention masks
17def fix_processor_outputs(inputs):
18 if 'cross_attention_mask' in inputs and 0 in inputs['cross_attention_mask'].shape:
19 batch_size, seq_len, _, num_tiles = inputs['cross_attention_mask'].shape
20 visual_features = 6404 # Critical dimension
21 new_mask = torch.ones((batch_size, seq_len, visual_features, num_tiles),
22 device=inputs['cross_attention_mask'].device)
23 inputs['cross_attention_mask'] = new_mask
24 return inputs
25
26# Function to process multiple images
27def process_multiple_images(original_image, cam_image, cam_overlay, comparison_image, query):
28 # Process with all four images
29 # Note: This is a simplified approach and may need adaptation based on model capabilities
30 inputs = processor(
31 images=[original_image, cam_image, cam_overlay, comparison_image],
32 text=query,
33 return_tensors="pt"
34 )
35
36 # Fix cross-attention mask
37 inputs = fix_processor_outputs(inputs)
38 inputs = {k: v.to(model.device) for k, v in inputs.items() if isinstance(v, torch.Tensor)}
39
40 # Generate output
41 with torch.no_grad():
42 output_ids = model.generate(**inputs, max_new_tokens=500)
43 response = processor.decode(output_ids[0], skip_special_tokens=True)
44 return response
45
46# Example usage
47original_image = Image.open("path/to/original.jpg").convert("RGB")
48cam_image = Image.open("path/to/cam.jpg").convert("RGB")
49cam_overlay = Image.open("path/to/overlay.jpg").convert("RGB")
50comparison_image = Image.open("path/to/comparison.jpg").convert("RGB")
51
52query = "Analyze these images and explain if they show a deepfake."
53response = process_multiple_images(original_image, cam_image, cam_overlay, comparison_image, query)
54print(response)