-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
133 lines (110 loc) · 5.07 KB
/
Copy pathtrain.py
File metadata and controls
133 lines (110 loc) · 5.07 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
import os
import argparse
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
from dz_tdpo.config import TDPODKLConfig
from dz_tdpo.model import TemporalCausalLM
from dz_tdpo.trainer import TDPODKLTrainer
from dz_tdpo.loss import SimPOLoss, TDPO_DKLLoss
from dz_tdpo.data.dataset import TemporalPreferenceDataset
from dz_tdpo.data.msc import msc_to_temporal_preference
from dz_tdpo.data.ultrachat import build_ultrachat_dataset
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--data_dir", type=str, required=True, help="Path to MSC/UltraChat dataset folder")
parser.add_argument("--model_name_or_path", type=str, default="microsoft/Phi-3.5-mini-instruct")
parser.add_argument("--output_dir", type=str, default="./checkpoints")
parser.add_argument("--loss_type", type=str, default="tdpo", choices=["tdpo", "simpo"])
parser.add_argument("--epochs", type=int, default=4)
parser.add_argument("--batch_size", type=int, default=2)
parser.add_argument("--lr", type=float, default=1.5e-5)
parser.add_argument("--use_temporal_bias", action="store_true")
parser.add_argument("--use_adaptive_tau", action="store_true")
return parser.parse_args()
def main():
args = parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if device.type == 'cuda':
print(f"CUDA detected, using GPU: {torch.cuda.get_device_name(0)}")
print(f"Current GPU video memory: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.2f} GB")
else:
print("CUDA not detected, using CPU training")
config = TDPODKLConfig(
model_name=args.model_name_or_path,
tau=args.tau,
beta0=args.beta0,
use_temporal_bias=args.use_temporal_bias,
loss_type=args.loss_type
)
tokenizer = AutoTokenizer.from_pretrained(args.model_name_or_path)
if tokenizer.eos_token is None:
tokenizer.add_special_tokens({'eos_token': '<|end|>'})
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
special_tokens_dict = {'additional_special_tokens': ["<|user|>", "<|assistant|>", "<|end|>"]}
num_added_toks = tokenizer.add_special_tokens(special_tokens_dict)
tokenizer.eos_token_id = tokenizer.convert_tokens_to_ids("<|end|>")
print(f"Token increase the quantity: {num_added_toks}")
model_dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32
print(f"Load the model to the device, using dtype: {model_dtype}...")
base_model = AutoModelForCausalLM.from_pretrained(
args.model_name_or_path,
dtype=model_dtype,
attn_implementation="sdpa"
).to(device)
base_model.gradient_checkpointing_enable()
policy_model = TemporalCausalLM(base_model, config, device)
ref_model = None
if args.loss_type == "tdpo":
ref_base = AutoModelForCausalLM.from_pretrained(
args.model_name_or_path,
dtype=model_dtype,
attn_implementation="sdpa"
).to(device)
if num_added_toks > 0:
print(f"Resizing model embeddings to {len(tokenizer)}")
base_model.resize_token_embeddings(len(tokenizer))
ref_base.resize_token_embeddings(len(tokenizer))
ref_config = TDPODKLConfig(model_name=args.model_name_or_path, use_temporal_bias=False)
ref_model = TemporalCausalLM(ref_base, ref_config, device)
print("Loading MSC training set...")
raw_train_ds = msc_to_temporal_preference(
tokenizer,
data_dir=args.data_dir,
split="train",
sessions_per_sample=4,
neg_distance=5
)
raw_val_ds = msc_to_temporal_preference(
tokenizer,
data_dir=args.data_dir,
split="validation",
sessions_per_sample=4,
neg_distance=5
)
train_dataset = TemporalPreferenceDataset(raw_train_ds.samples, tokenizer, config)
val_dataset = TemporalPreferenceDataset(raw_val_ds.samples, tokenizer, config)
print(f"MSC training samples:{len(train_dataset)}")
print(f"MSC validation sample:{len(val_dataset)}")
for i in range(min(5, len(train_dataset))):
sample = train_dataset[i]
for key in ['input_ids', 'chosen_reply_ids', 'rejected_reply_ids']:
tensor = getattr(sample, key, None)
if tensor is not None and tensor.max() >= tokenizer.vocab_size:
print(f"sample {i} is {key} cross-border: {tensor.max()} >= {tokenizer.vocab_size}")
trainer = TDPODKLTrainer(
policy_model=policy_model,
ref_model=ref_model,
tokenizer=tokenizer,
config=config,
train_dataset=train_dataset,
val_dataset=val_dataset,
learning_rate=1.5e-5,
device=device,
gradient_accumulation_steps=8,
output_dir=args.output_dir
)
trainer.train(num_epochs=args.epochs)
trainer.save_checkpoint("final_model.pt")
if __name__ == "__main__":
main()