Repository navigation
Expand file tree
/
Copy pathutils.py
More file actions
152 lines (126 loc) · 6.36 KB
/
Copy pathutils.py
File metadata and controls
152 lines (126 loc) · 6.36 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
import copy
import torch
import argparse
import clip
import json
from pathlib import Path
def set_seed(seed):
"""Set random seed for reproducibility."""
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
import random, numpy as np
random.seed(seed)
np.random.seed(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def freeze_model(model):
"""Freeze all parameters of a model."""
for p in model.parameters():
p.requires_grad = False
def copy_model_and_freeze(accelerator, model_to_copy):
"""Return UNWRAPPED copied model"""
unwrapped_model = accelerator.unwrap_model(model_to_copy)
ref_model = copy.deepcopy(unwrapped_model)
ref_model = ref_model.to(accelerator.device)
freeze_model(ref_model)
ref_model.eval()
return ref_model
def update_copied_model_param(accelerator, model_to_copy, copied_model_vault):
unwrapped_model = accelerator.unwrap_model(model_to_copy)
with torch.no_grad():
for ref_param, student_param in zip(unwrapped_model.parameters(), copied_model_vault.parameters()):
# 使用 .copy_() 直接在预先分配的显存上硬覆写,速度极快
ref_param.data.copy_(student_param.data)
DATASET_NAMES = [
"aircraft", "caltech101", "cifar10", "cifar100", "dtd",
"eurosat", "flowers", "food", "mnist", "oxford_pet", "stanford_cars",
]
def parse_args():
parser = argparse.ArgumentParser(description="ZSCL Training Framework")
# Model
parser.add_argument("--model", type=str, default="ViT-B/16",
choices=clip.available_models(),
help="CLIP model architecture")
# Dataset
parser.add_argument("--datasets", type=list[str], default=[
"aircraft", "caltech101", "cifar100", "dtd",
"eurosat", "flowers", "food", "mnist", "oxford_pet", "stanford_cars",
],
nargs="+",
choices=DATASET_NAMES,
help="Target downstream datasets for continual training")
parser.add_argument("--data_root", type=str, default="./data",
help="Root directory for all datasets")
parser.add_argument("--batch_size", type=int, default=32)
parser.add_argument("--batch_size_eval", type=int, default=32)
parser.add_argument("--num_workers", type=int, default=16)
parser.add_argument("--grad_accumulation", type=int, default=1)
parser.add_argument("--few_shot", type=int, default=None)
# Training
parser.add_argument("--train_count", type=str, default='iter', choices=['iter', 'epoch'])
parser.add_argument("--iterations", type=int, default=1000)
parser.add_argument("--epochs", type=int, default=None)
parser.add_argument("--merge_freq", type=int, default=100)
parser.add_argument("--lr", type=float, default=1e-5,
help="Peak learning rate")
parser.add_argument("--min_lr", type=float, default=0.0,
help="Minimum learning rate for cosine decay")
parser.add_argument("--weight_decay", type=float, default=0.01)
parser.add_argument("--warmup_ratio", type=float, default=0.1,
help="Fraction of total steps used for LR warmup")
parser.add_argument("--max_grad_norm", type=float, default=1.0,
help="Max gradient norm for clipping (0 = disabled)")
# Loss weights
parser.add_argument("--ce_weight", type=float, default=1.0,
help="Weight for cross-entropy classification loss")
parser.add_argument("--l2_wc_weight", type=float, default=1.0,
help="Weight for L2-WC loss")
parser.add_argument("--logit_distill_weight", type=float, default=1,
help="Weight for logit-level distillation loss, lambda (0 = disabled)")
parser.add_argument("--distill_temperature", type=float, default=2.0,
help="Temperature for logit distillation")
# Finetuning strategy
parser.add_argument("--freeze_text_encoder", action="store_true", default=False,
help="Freeze the text encoder during finetuning")
parser.add_argument("--freeze_visual_layers", type=int, default=0,
help="Number of visual transformer layers to freeze (from bottom)")
# Logging & Checkpointing
parser.add_argument("--output_dir", type=str, default="./output",
help="Directory to save checkpoints and logs")
parser.add_argument("--eval_freq", type=int, default=1,
help="Evaluate every N epochs")
parser.add_argument("--save_freq", type=int, default=5,
help="Save checkpoint every N epochs")
parser.add_argument("--eval_datasets", type=str, nargs="*", default=None,
help="Additional datasets to evaluate zero-shot transfer on "
"(e.g., --eval_datasets cifar10 dtd)")
# Misc
parser.add_argument("--seed", type=int, default=42)
return parser.parse_args()
class RunManager:
"""Manages the output directory, saving arguments, checkpoints, and history."""
def __init__(self, base_output_dir, run_name, args):
self.dir = Path(base_output_dir) / run_name
self.dir.mkdir(parents=True, exist_ok=True)
self.args = args
self._save_args()
def _save_args(self):
with open(self.dir / "args.json", "w", encoding="utf-8") as f:
json.dump(vars(self.args), f, indent=2)
def save_model(self, filename, epoch, model_state_dict, **extra_data):
data = {
"epoch": epoch,
"model_state_dict": model_state_dict,
"args": vars(self.args)
}
data.update(extra_data)
torch.save(data, self.dir / filename)
def save_best_model(self, epoch, model_state_dict, accuracy):
self.save_model("best_model.pt", epoch, model_state_dict, accuracy=accuracy)
def save_checkpoint(self, epoch, model_state_dict):
self.save_model(f"checkpoint_epoch{epoch}.pt", epoch, model_state_dict)
def save_final_model(self, epoch, model_state_dict):
self.save_model("final_model.pt", epoch, model_state_dict)
def save_history(self, history):
with open(self.dir / "history.json", "w", encoding="utf-8") as f:
json.dump(history, f, indent=2)