-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval.py
More file actions
256 lines (228 loc) · 13.3 KB
/
Copy patheval.py
File metadata and controls
256 lines (228 loc) · 13.3 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
import os
import gc
import torch
from functools import partial
from typing import Optional, Tuple
from dataclasses import dataclass, field
from transformers import (
AutoTokenizer, DeformableDetrConfig,
HfArgumentParser, TrainingArguments, Trainer
)
from loader import DVCDataset, trainer_collate_fn, get_loader
from pdvc import DeformableDetrForObjectDetection
from captioners import LSTMCaptioner, MBartDecoderCaptioner
from evaluation import preprocess_logits_for_metrics, compute_metrics
from config import TGT_LANG, TRIMMED_TOKENIZER_DIR
from test import debug_variance, debug_encoder_layers, debug_query_variance
from config import *
@dataclass
class ModelArguments:
d_model: int = field(default=1024)
encoder_layers: int = field(default=2)
decoder_layers: int = field(default=2)
encoder_attention_heads: int = field(default=8)
decoder_attention_heads: int = field(default=8)
encoder_n_points: int = field(default=4)
decoder_n_points: int = field(default=4)
num_feature_levels: int = field(default=4, metadata={'help': 'The number of input feature levels'})
num_queries: int = field(default=30, metadata={'help': 'Maximum number of events a window can have'})
num_labels: int = field(default=1, metadata={'help': 'Single foreground class for caption'})
auxiliary_loss: bool = field(default=True, metadata={'help': 'The training step may spend a time in per-layer caption alignment and Hungarian matching'})
class_cost: float = field(default=2, metadata={'help': 'LOSS weight of the classification (focal) term'})
bbox_cost: float = field(default=0, metadata={'help': 'LOSS weight of the L1 box term (0 = disabled, matches paper L_total)'})
giou_cost: float = field(default=4, metadata={'help': 'LOSS weight of the GIoU term'})
counter_cost: float = field(default=2, metadata={'help': 'LOSS weight of the event counter (BCE) term'})
caption_cost: float = field(default=2, metadata={'help': 'LOSS weight of the captioning (NLL) term'})
# Hungarian MATCHER cost weights (App C.2): (cls, L1, giou) = (1, 5, 2), distinct from loss weights.
match_class_cost: float = field(default=1, metadata={'help': 'MATCHER classification cost weight'})
match_bbox_cost: float = field(default=5, metadata={'help': 'MATCHER L1 cost weight (dominant)'})
match_giou_cost: float = field(default=2, metadata={'help': 'MATCHER GIoU cost weight'})
focal_alpha: float = field(default=0.25)
with_box_refine: bool = field(default=True, metadata={'help': 'Learnt (True) or Ground truth proposals (False, all losses except caption loss will be disabled)'})
# Caption head / decoder bits
num_cap_layers: int = field(default=3)
cap_dropout_rate: float = field(default=0.1)
captioner_type: str = field(default='mbart', metadata={'help': 'Type of captioner to use (mbart or lstms)'})
@dataclass
class DataArguments:
max_tries: int = field(default=20, metadata={'help': 'Maximum attempts to find a valid window with at least one event'})
noise_rate: float = field(default=0.15, metadata={'help': 'Proportion of words to mask for noise injection during non-streaming training'})
pose_augment: bool = field(default=False, metadata={'help': 'Apply pose augmentation during training'})
stride_ratio: float = field(default=0.9, metadata={'help': 'Stride ratio for window sampling during validation/testing'})
min_events: int = field(default=1, metadata={'help': 'Minimum number of events in a window'})
max_events: int = field(default=10, metadata={'help': 'Maximum number of events in a window'})
max_event_tokens: int = field(default=40, metadata={'help': 'Maximum number of tokens per event/caption'})
max_window_tokens: int = field(default=128, metadata={'help': 'Maximum number of tokens in a window for non-streaming input'})
load_by: str = field(default='window', metadata={'help': "Load data by 'window' or by 'video'"})
# Metrics/Ranking
ranking_temperature: float = field(default=2.0, metadata={'help': 'Exponent T in caption score normalization by length^T'})
alpha: float = field(default=0.3, metadata={'help': 'Ranking policy: joint_score = alpha * (caption_score / len(tokens)^T) + (1 - alpha) * det_score'})
top_k: int = field(default=20, metadata={'help': 'Keep top k events during evaluation for metrics computation'})
temporal_iou_thresholds: Tuple[float, float, float, float] = field(default=(0.3, 0.5, 0.7, 0.9))
soda_recursion_limit: int = field(default=0, metadata={'help': 'Increase recursion limit for SODA_c DP if needed, 0 to disable for faster calculations'})
aggregation_mode: str = field(default='video', metadata={
'help': f"How to aggregate caption pairs into one score per IoU threshold. Same metric KEY NAMES across modes; only internal compute changes. Options: "
f"'corpus' (ActivityNet default; pool ALL pairs -> sacrebleu corpus, geometric-mean over n-grams, non-linear when combining splits); "
f"'window' (corpus-text-metrics per window then mean; linear-aggregating); "
f"'video' (group windows by video_id, corpus per video, then mean across videos)."})
@dataclass
class EvalArguments:
checkpoint_path: str = field(default=CHECKPOINT_DIR, metadata={'help': 'Path to the checkpoint directory to evaluate'})
output_dir: str = field(default='/tmp/eval', metadata={'help': 'Directory for evaluation outputs'})
per_device_eval_batch_size: int = field(default=32, metadata={'help': 'Can be higher if greedy but should be smaller if using beam search'})
dataloader_num_workers: int = field(default=4, metadata={'help': 'Number of subprocesses to use for data loading'})
seed: int = field(default=42, metadata={'help': 'Random seed for reproducibility'})
fp16: bool = field(default=False, metadata={'help': 'Use mixed precision evaluation'})
bf16: bool = field(default=False, metadata={'help': 'Use bfloat16 for mixed precision evaluation'})
eval_val: bool = field(default=True, metadata={'help': 'Evaluate on validation set'})
eval_test: bool = field(default=True, metadata={'help': 'Evaluate on test set'})
def main():
# Parse CLI args
parser = HfArgumentParser((ModelArguments, DataArguments, EvalArguments))
model_args, data_args, eval_args = parser.parse_args_into_dataclasses()
device = 'cuda' if torch.cuda.is_available() else 'cpu'
print(f'Loading checkpoint from: {eval_args.checkpoint_path}')
# Model Setup
tokenizer = AutoTokenizer.from_pretrained(TRIMMED_TOKENIZER_DIR)
config = DeformableDetrConfig(
d_model=model_args.d_model,
encoder_layers=model_args.encoder_layers,
decoder_layers=model_args.decoder_layers,
encoder_attention_heads=model_args.encoder_attention_heads,
decoder_attention_heads=model_args.decoder_attention_heads,
encoder_n_points=model_args.encoder_n_points,
decoder_n_points=model_args.decoder_n_points,
activation_function='gelu',
num_feature_levels=model_args.num_feature_levels,
num_queries=model_args.num_queries,
num_labels=model_args.num_labels,
auxiliary_loss=model_args.auxiliary_loss,
# config.{class,bbox,giou}_cost feed the Hungarian MATCHER — use matching weights (1, 5, 2).
class_cost=model_args.match_class_cost,
bbox_cost=model_args.match_bbox_cost,
giou_cost=model_args.match_giou_cost,
focal_alpha=model_args.focal_alpha,
with_box_refine=model_args.with_box_refine,
)
model = DeformableDetrForObjectDetection(
config=config,
captioner_class=MBartDecoderCaptioner if model_args.captioner_type=='mbart' else LSTMCaptioner,
vocab_size=tokenizer.vocab_size,
bos_token_id=tokenizer.bos_token_id,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.pad_token_id,
decoder_start_token_id=tokenizer.lang_code_to_id[TGT_LANG],
num_cap_layers=model_args.num_cap_layers,
cap_dropout_rate=model_args.cap_dropout_rate,
max_event_tokens=data_args.max_event_tokens,
max_events=data_args.max_events,
use_gt_boxes_for_caption=not model_args.with_box_refine, # No GT boxes needed in 2-stage training
weight_dict={
'loss_ce': model_args.class_cost, 'loss_bbox': model_args.bbox_cost, 'loss_giou': model_args.giou_cost,
'loss_counter': model_args.counter_cost, 'loss_caption': model_args.caption_cost
}
).to(device)
model.load_state_dict(torch.load(os.path.join(eval_args.checkpoint_path, 'pytorch_model.bin')))
total_params = sum(p.numel() for p in model.parameters())
print(f'Model loaded with {total_params / 1e6:.2f}M parameters')
# Test model variance
if model_args.with_box_refine:
val_loader = get_loader(split='val', tokenizer=tokenizer, batch_size=4, max_events=10, max_event_tokens=40)
test_batch = next(iter(val_loader))
test_pixel_values = test_batch['pixel_values'].to(device)
test_pixel_mask = test_batch['pixel_mask'].to(device)
print('\n[Model variance before eval mode]')
debug_variance(model, test_pixel_values, test_pixel_mask)
debug_encoder_layers(model, test_pixel_values, test_pixel_mask)
debug_query_variance(model, test_pixel_values, test_pixel_mask)
model.eval() # Test after enabling inference mode
if model_args.with_box_refine:
print('\n[Model variance after eval mode]')
debug_variance(model, test_pixel_values, test_pixel_mask)
debug_encoder_layers(model, test_pixel_values, test_pixel_mask)
debug_query_variance(model, test_pixel_values, test_pixel_mask)
# Evaluation
training_args = TrainingArguments( # Create training args for Trainer (required even for evaluation)
output_dir=eval_args.output_dir,
per_device_eval_batch_size=eval_args.per_device_eval_batch_size,
dataloader_num_workers=eval_args.dataloader_num_workers,
seed=eval_args.seed,
fp16=eval_args.fp16,
bf16=eval_args.bf16,
report_to='none',
do_train=False,
do_eval=True,
)
if eval_args.eval_val:
print('\n' + '='*80)
print('Evaluating on validation set...')
print('='*80)
val_dataset = DVCDataset(
split='val', tokenizer=tokenizer, pose_augment=False, stride_ratio=data_args.stride_ratio,
min_events=data_args.min_events, max_events=data_args.max_events, max_event_tokens=data_args.max_event_tokens,
max_window_tokens=data_args.max_window_tokens, load_by=data_args.load_by, seed=eval_args.seed
)
print(f'Val dataset: {len(val_dataset)} samples')
eval_trainer = Trainer(
model=model,
args=training_args,
eval_dataset=val_dataset,
data_collator=trainer_collate_fn,
preprocess_logits_for_metrics=preprocess_logits_for_metrics,
compute_metrics=partial(
compute_metrics,
ranking_temperature=data_args.ranking_temperature,
alpha=data_args.alpha,
top_k=data_args.top_k,
temporal_iou_thresholds=data_args.temporal_iou_thresholds,
tokenizer=tokenizer,
soda_recursion_limit=data_args.soda_recursion_limit,
aggregation_mode=data_args.aggregation_mode,
eval_windows=val_dataset.eval_windows,
),
)
val_metrics = eval_trainer.evaluate(metric_key_prefix='val')
print('\nValidation metrics:', val_metrics, sep='\n')
del val_dataset, eval_trainer
gc.collect()
if torch.cuda.is_available(): torch.cuda.empty_cache()
if eval_args.eval_test:
print('\n' + '='*80)
print('Evaluating on test set...')
print('='*80)
test_dataset = DVCDataset(
split='test', tokenizer=tokenizer, pose_augment=False, stride_ratio=data_args.stride_ratio,
max_event_tokens=data_args.max_event_tokens, max_window_tokens=data_args.max_window_tokens,
min_events=data_args.min_events, load_by=data_args.load_by, seed=eval_args.seed
)
print(f'Test dataset: {len(test_dataset)} samples')
eval_trainer = Trainer(
model=model,
args=training_args,
eval_dataset=test_dataset,
data_collator=trainer_collate_fn,
preprocess_logits_for_metrics=preprocess_logits_for_metrics,
compute_metrics=partial(
compute_metrics,
ranking_temperature=data_args.ranking_temperature,
alpha=data_args.alpha,
top_k=data_args.top_k,
temporal_iou_thresholds=data_args.temporal_iou_thresholds,
tokenizer=tokenizer,
soda_recursion_limit=data_args.soda_recursion_limit,
aggregation_mode=data_args.aggregation_mode,
eval_windows=test_dataset.eval_windows,
),
)
test_metrics = eval_trainer.evaluate(metric_key_prefix='test')
print('\nTest metrics:', test_metrics, sep='\n')
del test_dataset, eval_trainer
gc.collect()
if torch.cuda.is_available(): torch.cuda.empty_cache()
# Cleanup
model.to('cpu')
del tokenizer, model
gc.collect()
if torch.cuda.is_available(): torch.cuda.empty_cache()
if __name__ == '__main__':
main()