simplecloud/VidChain-exercise
✏️ Data for VidChain Excercise VidChain: Chain-of-Tasks with Metric-based Direct Preference Optimization for Dense Video Captioning Ji Soo Lee*, Jongha Kim*, Jeehye Na, Jinyoung Park, Hyunwoo J. Kim†. AAAI 2025 🎯 Learning Objectives By working through this exercise, you will: Reproduce baseline behavior of a video-language model (VTimeLLM, CVPR 2024 Highlight). Observe the limitations of existing approaches in temporal… See the full description on the dataset page: https://huggingface.co/datasets/simplecloud/VidChain-exercise.
0150
1from dvc_eval import eval_dvc, eval_soda2import json3import argparse4import re5import difflib6import os7from torchvision.transforms import Compose, Resize, CenterCrop, Normalize8import torch9 10# Define image transforms11try:12 from torchvision.transforms import InterpolationMode13 BICUBIC = InterpolationMode.BICUBIC14except ImportError:15 BICUBIC = Image.BICUBIC16 from torchvision.transforms import Compose, Resize, CenterCrop, Normalize17 18 19transform = Compose([20 Resize(224, interpolation=BICUBIC),21 CenterCrop(224),22 Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)),23])24 25# Check if model files exist26def check_model_files(config):27 """Check if required model files exist"""28 files_to_check = [29 config.clip_path,30 config.pretrain_mm_mlp_adapter,31 config.stage2,32 config.stage3,33 config.stage4,34 config.stage5,35 config.model_base36 ]37 38 missing_files = []39 for file_path in files_to_check:40 if not os.path.exists(file_path):41 missing_files.append(file_path)42 43 if missing_files:44 print("⚠ Missing model files:")45 for file_path in missing_files:46 print(f" - {file_path}")47 print("\nPlease download the required model checkpoints.")48 return False49 else:50 print("✓ All model files found")51 return True52 53# CLIP Utility Functions54 55 56# Video Utility Functions57# Utility functions for video processing58def extract_video_features(video_path, clip_model, video_loader, transform):59 """Extract features from a video file"""60 try:61 # Extract frames from video62 _, images = video_loader.extract({'id': None, 'video': video_path})63 64 # Apply transforms65 images = transform(images / 255.0)66 images = images.to(torch.float16)67 68 # Encode with CLIP69 with torch.no_grad():70 features = clip_model.encode_image(images.to('cuda'))71 72 return features73 except Exception as e:74 print(f"Error processing video {video_path}: {e}")75 return None76 77def find_video_file(video_id, video_folder):78 """Find video file with various extensions"""79 for ext in ['mp4', 'mkv', 'webm', 'avi', 'mov']:80 video_path = os.path.join(video_folder, f"{video_id}.{ext}")81 if os.path.isfile(video_path):82 return video_path83 return None84 85def load_dataset(data_path):86 """Load dataset from JSON file"""87 try:88 with open(data_path, 'r') as f:89 data = json.load(f)90 return data91 except Exception as e:92 print(f"✗ Error loading dataset: {e}")93 return None94 95 96## EVALUTE FUNCTIONS97def merge_similar_sentences(data):98 if not data: return data99 merged_data = []100 current_sentence = data[0]["sentence"]101 current_timestamp = data[0]["timestamp"]102 for i in range(1, len(data)):103 next_sentence = data[i]["sentence"]104 next_timestamp = data[i]["timestamp"]105 if difflib.SequenceMatcher(None, current_sentence, next_sentence).ratio() > 0.98 and -1 <= next_timestamp[0] - current_timestamp[1] <= 1:106 current_timestamp = [current_timestamp[0], next_timestamp[1]]107 else:108 merged_data.append({"sentence": current_sentence, "timestamp": current_timestamp})109 current_sentence = next_sentence110 current_timestamp = next_timestamp111 merged_data.append({"sentence": current_sentence, "timestamp": current_timestamp})112 return merged_data113 114def evaluate(id, event, timestamps, answer, js):115 pred = {}116 pred[id] = []117 for num in range(len(event)):118 pred[id].append({119 'timestamp': timestamps[num],120 'sentence': event[num]121 })122 123 refined_pred = []124 for num_pred, curr_pred in enumerate(pred[id]):125 duplicate = False126 for curr_pred2 in pred[id][num_pred + 1:]:127 128 if curr_pred2 == curr_pred:129 num_duplicates+=1130 duplicate=True131 132 133 if not duplicate:134 refined_pred.append(curr_pred)135 136 pred[id] = refined_pred137 gt_js = {k: v for k, v in js.items() if k in pred.keys()}138 139 for id, items in list(pred.items()): 140 items = merge_similar_sentences(items)141 duration = gt_js[id]['duration']142 for item in items:143 item['timestamp'][0] = item['timestamp'][0] * duration / 100144 item['timestamp'][1] = (item['timestamp'][1] + 1) * duration / 100145 pred[id] = items146 147 pred_result = {'results': pred}148 metrics = eval_soda(pred_result, [gt_js], print_matrix=False)149 metrics.update(eval_dvc(pred_result, [gt_js], 150 tious=[0.3, 0.5, 0.7], 151 distances=[],152 max_proposals_per_video=1000, 153 verbose=False, 154 no_lang_eval=False))155 156 print(f"Found {len(pred)} logs")157 metrics = {k: v.item() * 100 for k, v in metrics.items() if k in ['soda_c', 'METEOR', 'CIDEr']}158 return metrics