-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathmain.py
More file actions
155 lines (134 loc) · 6.48 KB
/
Copy pathmain.py
File metadata and controls
155 lines (134 loc) · 6.48 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
import argparse
import os
import numpy as np
import torch
import gc
from utils import load_model
from data_prepare import LoadDataset
from torch.utils.data import DataLoader, DistributedSampler
from torch.nn.parallel import DistributedDataParallel as DDP
from cal_KFAC import cal_ihvp, cal_grad, cal_influence
import warnings
import datetime
import pandas as pd
warnings.filterwarnings("ignore", category=FutureWarning, message=".*weights_only=False.*")
def setup_distributed():
torch.distributed.init_process_group(backend="nccl", timeout=datetime.timedelta(seconds=180000))
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
return local_rank
def main(args):
local_rank = setup_distributed()
# load model
model, tokenizer = load_model(args.model_path)
train_sample_rate = 0.5
val_sample_rate = 1.0
# calculate KFAC
train_dataset = LoadDataset(all_file_paths=args.full_train,
tokenizer=tokenizer,
max_seq_length=1024,
sample_percentage=train_sample_rate)
train_sampler = DistributedSampler(train_dataset, shuffle=True)
train_dataloader = DataLoader(train_dataset, sampler=train_sampler, batch_size=8, drop_last=True)
cal_ihvp(train_dataloader, model, args.save_path, args, local_rank)
del train_dataloader
del train_sampler
del train_dataset
torch.cuda.empty_cache()
gc.collect()
per_val_list = []
# calculate validation gradient
validation_dataset = LoadDataset(all_file_paths=args.validation_path,
tokenizer=tokenizer,
max_seq_length=1024,
sample_percentage=val_sample_rate)
validation_sampler = DistributedSampler(validation_dataset, shuffle=True)
validation_dataloader = DataLoader(validation_dataset, sampler=validation_sampler, batch_size=8)
validation_grad_path = cal_grad(validation_dataloader, model, args.save_path + f"/val_avg_grad", args, local_rank)
per_val_list.append(validation_grad_path)
del validation_dataloader
del validation_sampler
del validation_dataset
gc.collect()
for subset in list_subdirectories(args.validation_path):
subset_path = os.path.join(args.validation_path, subset)
if os.path.exists(subset_path):
dataset = LoadDataset(all_file_paths=subset_path,
tokenizer=tokenizer,
max_seq_length=1024,
sample_percentage=val_sample_rate)
sampler = DistributedSampler(dataset, shuffle=True)
dataloader = DataLoader(dataset, sampler=sampler, batch_size=8)
subset_grad_path = cal_grad(dataloader, model, args.save_path + f"/val_{subset}_grad", args, local_rank)
per_val_list.append(subset_grad_path)
# release memory
del dataloader
del sampler
del dataset
torch.cuda.empty_cache()
gc.collect()
if local_rank == 0:
print(per_val_list)
final_result = np.zeros((len(per_val_list), len(args.sub_train)))
per_train_list = []
for i, subset in enumerate(args.sub_train):
subset_path = os.path.join(args.full_train, subset)
if os.path.exists(subset_path):
dataset = LoadDataset(all_file_paths=subset_path,
tokenizer=tokenizer,
max_seq_length=1024,
sample_percentage=train_sample_rate)
sampler = DistributedSampler(dataset, shuffle=True)
dataloader = DataLoader(dataset, sampler=sampler, batch_size=8)
avg_grad_path = cal_grad(dataloader, model, args.save_path + f"/train_{subset}_grad", args, local_rank)
per_train_list.append(avg_grad_path)
# release memory
del dataloader
del sampler
del dataset
torch.cuda.empty_cache()
gc.collect()
if local_rank == 0:
print(per_train_list)
del model, tokenizer
torch.cuda.empty_cache()
gc.collect()
for i in range(len(per_train_list)):
for j in range(len(per_val_list)):
influence_list = cal_influence(hessian_path = args.save_path,
train_grad_path = per_train_list[i],
validation_grad_path = per_val_list[j],
local_rank = local_rank)
total_sum = sum(influence_list)
final_result[j, i] = total_sum.item()
del influence_list
torch.cuda.empty_cache()
gc.collect()
if local_rank == 0:
print(final_result)
df = pd.DataFrame(final_result, index=per_val_list, columns=per_train_list)
df.to_csv(args.save_path + f"/influence.csv", index=True, header=True)
def list_subdirectories(directory):
subdirectories = []
for root, dirs, files in os.walk(directory):
for dir_name in dirs:
subdirectories.append(dir_name)
return subdirectories
if __name__=='__main__':
parser = argparse.ArgumentParser()
parser.add_argument("--model-path", type=str, default=f"Path of YOUR base model")
parser.add_argument("--full-train", type=str, default=f"Path of YOUR full train dataset")
parser.add_argument("--validation-path", type=str, default=f"Path of YOUR validation dataset")
parser.add_argument("--sub-train", type=list, default=[f"Mathematics","Coding","bbh","Instruction","TrustAI"])
parser.add_argument("--save-path", type=str, default=f"Path of YOUR save folder")
parser.add_argument("--use-full-layer", type=bool, default=True)
parser.add_argument("--target-layers", type=list, default=["model.layers.1.mlp.gate_proj", "model.layers.5.mlp.gate_proj", "model.layers.10.mlp.gate_proj", "model.layers.15.mlp.gate_proj"
"model.layers.20.mlp.gate_proj", "model.layers.24.mlp.gate_proj", "model.layers.25.mlp.gate_proj", "model.layers.26.mlp.gate_proj", "model.layers.27.mlp.gate_proj", "model.layers.28.mlp.gate_proj"])
parser.add_argument("--without-output", type=bool, default=True)
parser.add_argument("--without-attention", type=bool, default=True)
args = parser.parse_args()
if not os.path.exists(args.save_path):
os.makedirs(args.save_path, exist_ok=True)
print(f"Directory '{args.save_path}' created.")
main(args)
torch.distributed.destroy_process_group()