-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathrun_vae_pretraining.py
More file actions
executable file
·266 lines (233 loc) · 10 KB
/
Copy pathrun_vae_pretraining.py
File metadata and controls
executable file
·266 lines (233 loc) · 10 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
257
258
259
260
261
262
263
264
265
266
#!/usr/bin/env python
# Copyright (c) 2018 The Google AI Language Team Authors and The HuggingFace Inc. team.
# Copyright (c) 2018 NVIDIA CORPORATION.
# Copyright (c) 2020 Microsoft Research.
# Copyright (c) 2025 ByteDance Ltd. and/or its affiliates.
# SPDX-License-Identifier: Apache-2.0
#
# This file has been modified by ByteDance Ltd. and/or its affiliates. on 2025
#
# Original file was released under Apache-2.0, with the full license text
# available at https://github.com/ChunyuanLI/Optimus/blob/master/code/examples/big_ae/run_lm_ae_pretraining.py.
#
# This modified file is released under the same license.
import os
import json
import logging
from pathlib import Path
from typing import Optional, Dict
from dataclasses import dataclass, field
from dataclasses_json import dataclass_json
import transformers
from models.trainer import CyclicalBetaTrainer
from models.auto_encoder import (
DATA_CONFIG_NAME,
MODEL_CONFIG_NAME,
TRAINER_CONFIG_NAME,
prepare_vae_model,
)
from utils.common import setup_logger, is_main_process
from utils.dataset import SupervisedLMDataset, LazySupervisedLMDataset
logger = logging.getLogger(__name__)
setup_logger(logger)
local_rank = None
@dataclass_json
@dataclass
class ModelArguments:
encoder_name_or_path: Optional[str] = field(
default="bert-base-cased", metadata={"help": "The encoder model checkpoint for weights initialization."}
)
decoder_name_or_path: Optional[str] = field(
default="gpt2", metadata={"help": "The decoder model checkpoint for weights initialization."}
)
encoder_max_length: int = field(
default=512, metadata={"help": "Expanded maximum sequence length. Sequences will be right padded (and possibly truncated)."},
)
decoder_max_length: int = field(
default=512, metadata={"help": "Expanded maximum sequence length. Sequences will be right padded (and possibly truncated)."},
)
freeze_decoder: Optional[bool] = field(
default=False, metadata={"help": "Whether to freeze decoder model paramenters."},
)
vae_latent_size: Optional[int] = field(
default=32, metadata={"help": "Latent space dimension."}
)
vae_adapter_size: Optional[int] = field(
default=8, metadata={"help": "Token numbers of adapter for both kv_memory and soft prompt."}
)
vae_latent_method: str = field(
default="soft_prompt",
metadata={
"help": "soft_prompt: VAE latent vector as a soft promt token embeddings (adding before the bos token); kv_memory: VAE latent vector as a past_key_values as previous memory; input_embed: VAE latent vector as a extra token embedding.",
"choices": ["soft_prompt", "kv_memory", "input_embed", "prefix_soft_prompt"],
},
)
vae_prefix_text: Optional[str] = field(
default=None, metadata={"help": "Init text for M+N soft prompts -> vae_latent_method==prefix_soft_prompt, set None to random init the M embeddings"}
)
threshold_kl: float = field(
default=None, metadata={"help": "The thresholding objective causes learning to give up driving down KL for dimensions of z that are already beneath the target compression rate."},
)
deterministic_connect: bool = field(
default=False, metadata={"help": "Use deterministic inference to generate latent codes, i.e., standard auto-encoders."},
)
length_weighted_loss: bool = field(
default=False, metadata={"help": "Whether to use length weighted loss."},
)
left_padding: bool = field(
default=False, metadata={"help": "Whether to use left padding for decoder. NOTE: it might have some issues with soft prompt and kv memory, see https://github.com/huggingface/peft/issues/1093 for detailed comparison."},
)
attn_implementation: str = field(
default="eager",
metadata={
"help": "eager: manual implementation of the attention; sdpa(torch>=2.1.1): attention using torch.nn.functional.scaled_dot_product_attention; or flash_attention_2: attention using Dao-AILab/flash-attention",
"choices": ["eager", "sdpa", "flash_attention_2"],
},
)
@dataclass_json
@dataclass
class DataArguments:
data_path: str = field(
default=None, metadata={"help": "Path to the training data."}
)
eval_data_path: str = field(
default=None, metadata={"help": "Path to the training data."}
)
lazy_preprocess: bool = field(
default=False, metadata={"help": "Whether to lazy load dataset."}
)
mlm: bool = field(
default=False, metadata={"help": "Train with masked-language modeling loss instead of language modeling."},
)
mlm_probability: float = field(
default=0.15, metadata={"help": "Ratio of tokens to mask for masked language modeling loss."},
)
shuffle_column: bool = field(
default=False, metadata={"help": "Train with masked-language modeling loss instead of language modeling."},
)
@dataclass_json
@dataclass
class TrainingArguments(transformers.TrainingArguments):
local_rank: int = field(default=-1, metadata={"help": "For distributed training: local_rank"})
cache_dir: Optional[str] = field(default=None)
optim: str = field(default="adamw_torch")
save_best_model_at_end: bool = field(
default=False, metadata={"help": "Use cyclical target beta."},
)
do_cyclical: bool = field(
default=False, metadata={"help": "Use cyclical target beta."},
)
num_cycle: int = field(
default=10, metadata={"help": "The number of cyclical beta iterations."},
)
beta: float = field(
default=1.0, metadata={"help": "The weighting hyper-parameter of the KL term in VAE."},
)
ratio_zero: float = field(
default=0.5, metadata={"help": "Learning schedule, the percentage for the pure auto-encoding stage."},
)
ratio_increase: float = field(
default=0.25, metadata={"help": "Learning schedule, the percentage for the annealing stage."},
)
trust_remote_code: bool = field(
default=False, metadata={"help": "Force run some model when transformers version is lower than its requirments"},
)
def extract_losses(eval_preds):
loss_kl, loss_rec = eval_preds.predictions
return {
"loss_kl": loss_kl.mean(),
"loss_rec": loss_rec.mean(),
}
def make_supervised_mlm_data_module(
encoder_tokenizer: transformers.PreTrainedTokenizer,
decoder_tokenizer: transformers.PreTrainedTokenizer,
data_args,
) -> Dict:
if is_main_process(local_rank):
logger.info(f"[Data] Loading data from {data_args}.")
dataset_cls = (
LazySupervisedLMDataset if data_args.lazy_preprocess else SupervisedLMDataset
)
def load_json_dataset(file_path):
with open(file_path, "r") as rf:
if file_path.endswith("jsonl"):
res_json = [json.loads(line) for line in rf if line.strip() != ""]
elif file_path.endswith("json"):
res_json = json.load(rf)
else:
raise NotImplementedError(f"not supported file type -> {file_path}")
assert isinstance(res_json, list)
return res_json
train_json = load_json_dataset(data_args.data_path)
train_dataset = dataset_cls(
train_json,
encoder_tokenizer,
decoder_tokenizer,
use_mlm=data_args.mlm,
mlm_probability=data_args.mlm_probability,
)
if data_args.eval_data_path:
eval_json = load_json_dataset(data_args.eval_data_path)
eval_dataset = dataset_cls(
eval_json,
encoder_tokenizer,
decoder_tokenizer,
shuffle_column=data_args.shuffle_column,
use_mlm=data_args.mlm,
mlm_probability=data_args.mlm_probability,
)
else:
eval_dataset = None
print(train_dataset.source_texts[:5])
return dict(train_dataset=train_dataset, eval_dataset=eval_dataset)
def train(
trainer_cls: transformers.Trainer = CyclicalBetaTrainer,
train_arg_cls: transformers.TrainingArguments = TrainingArguments
):
global local_rank
parser = transformers.HfArgumentParser(
(ModelArguments, DataArguments, train_arg_cls)
)
model_args, data_args, training_args = parser.parse_args_into_dataclasses()
local_rank = training_args.local_rank
# reserve soft prompt tokens
if model_args.vae_latent_method == "soft_prompt":
model_args.decoder_max_length -= model_args.vae_adapter_size
if model_args.vae_latent_method == "prefix_soft_prompt":
model_args.decoder_max_length -= model_args.vae_adapter_size * 2
vae_model, encoder_tokenizer, decoder_tokenizer = prepare_vae_model(model_args, training_args, local_rank)
data_module = make_supervised_mlm_data_module(encoder_tokenizer, decoder_tokenizer, data_args)
# Start training
trainer = trainer_cls(
args=training_args,
model=vae_model,
tokenizer=encoder_tokenizer,
compute_metrics=extract_losses,
**data_module,
)
# save model / data / training arguments
if trainer.is_world_process_zero():
os.makedirs(training_args.output_dir, exist_ok=True)
for fname, _args in [
(MODEL_CONFIG_NAME, model_args),
(DATA_CONFIG_NAME, data_args),
(TRAINER_CONFIG_NAME, training_args),
]:
with open(os.path.join(training_args.output_dir, fname), "w") as wf:
json.dump(_args.to_dict(), wf, indent=4, ensure_ascii=False)
if list(Path(training_args.output_dir).glob("checkpoint-*")):
trainer.train(resume_from_checkpoint=True)
else:
trainer.train()
# Move the best model to output_dir
if training_args.save_best_model_at_end:
if trainer.is_world_process_zero():
best_ckpt_path = Path(trainer.state.best_model_checkpoint)
out_ckpt_path = Path(training_args.output_dir)
for each_file in best_ckpt_path.glob("*"): # grabs all files
each_file.rename(out_ckpt_path.joinpath(each_file.name)) # moves to output folder.
else:
trainer.save_state()
trainer.save_model()
if __name__ == "__main__":
train()