| import os |
| import json |
| import torch |
| import wandb |
| from datasets import load_from_disk |
| from transformers import ( |
| Trainer, |
| TrainingArguments, |
| DefaultDataCollator, |
| set_seed, |
| get_cosine_schedule_with_warmup, |
| AutoModelForCausalLM, AutoTokenizer,AutoConfig |
| ) |
| from torch.optim import AdamW |
|
|
| |
| |
| |
| |
| |
|
|
| MODEL_ID = "phasorkinetics/pmnet" |
| CONFIG_PATH = "/workspace/config_pmnet_no_mem.json" |
| CKPT_SAVE_DIR = "/workspace/ckpt/pmnet_no_mem" |
| CKPT_START_DIR = "/workspace/ckpt_start_no_mem" |
| compile_model = True |
|
|
| |
|
|
| DATA_DIR = "/data/copypaste_2048" |
| SEED = 42 |
|
|
| BATCH_SIZE = 256 |
| GRADIENT_ACCUMULATION_STEPS = 1 |
| NUM_DEVICES = 2 |
|
|
| os.environ["WANDB_PROJECT"] = "pmnet_copy_paste" |
| os.environ["WANDB_Run_ID"] = "" |
| os.environ["WANDB_RESUME"] = "allow" |
| |
|
|
| per_device_batch_size = BATCH_SIZE // (NUM_DEVICES*GRADIENT_ACCUMULATION_STEPS) |
|
|
|
|
| def get_optimizer_grouped_parameters(model, weight_decay: float): |
| no_decay_keywords = ["bias", "norm", "embedding", "layernorm", "a_log"] |
| decay_params = [] |
| no_decay_params = [] |
|
|
| for name, param in model.named_parameters(): |
| if not param.requires_grad: continue |
| if any(keyword in name.lower() for keyword in no_decay_keywords): |
| no_decay_params.append(param) |
| else: |
| decay_params.append(param) |
| return [ |
| {"params": decay_params, "weight_decay": weight_decay}, |
| {"params": no_decay_params, "weight_decay": 0.0}, |
| ] |
|
|
| def main(): |
| set_seed(SEED) |
|
|
| |
| config = AutoConfig.from_pretrained(MODEL_ID, trust_remote_code=True) |
| if CONFIG_PATH and os.path.exists(CONFIG_PATH): |
| with open(CONFIG_PATH, "r", encoding="utf-8") as f: |
| local_config_dict = json.load(f) |
| config.update(local_config_dict) |
| model = AutoModelForCausalLM.from_config(config, trust_remote_code=True) |
|
|
| |
| total_params = sum(p.numel() for p in model.parameters()) |
| print(f"Total params: {total_params/1_000_000}") |
|
|
| dataset_train = load_from_disk(os.path.join(DATA_DIR, "train")) |
| dataset_val = load_from_disk(os.path.join(DATA_DIR, "val")) |
|
|
| tokenizer = AutoTokenizer.from_pretrained("google/byt5-small") |
|
|
| training_args = TrainingArguments( |
| output_dir=CKPT_SAVE_DIR, |
| overwrite_output_dir=False, |
| num_train_epochs=10, |
|
|
| save_strategy="steps", |
| save_steps=100, |
| eval_strategy="steps", |
| eval_steps=100, |
| logging_strategy="steps", |
| logging_steps=10, |
| save_total_limit=2, |
|
|
| load_best_model_at_end=True, |
| metric_for_best_model="loss", |
| greater_is_better=False, |
|
|
| seed=SEED, |
| data_seed=SEED, |
|
|
| per_device_train_batch_size=per_device_batch_size, |
| per_device_eval_batch_size=per_device_batch_size, |
| gradient_accumulation_steps=GRADIENT_ACCUMULATION_STEPS, |
|
|
| learning_rate=5e-4, |
| max_grad_norm=1.0, |
| fp16=False, |
| bf16=True, |
|
|
| dataloader_num_workers=8, |
| report_to="wandb", |
| ddp_find_unused_parameters=False, |
|
|
| lr_scheduler_type="cosine", |
| warmup_steps=100, |
| torch_compile=compile_model |
| ) |
|
|
| data_collator = DefaultDataCollator() |
|
|
| optimizer_grouped_parameters = get_optimizer_grouped_parameters(model, weight_decay=0.1) |
| optimizer = AdamW(optimizer_grouped_parameters, lr=training_args.learning_rate, betas=(0.9, 0.95)) |
|
|
| trainer = Trainer( |
| model=model, |
| args=training_args, |
| train_dataset=dataset_train, |
| eval_dataset=dataset_val, |
| data_collator=data_collator, |
| optimizers=(optimizer, None), |
| ) |
| trainer.train(resume_from_checkpoint=CKPT_START_DIR if os.path.exists(CKPT_START_DIR) else None) |
| trainer.save_model(os.path.join(CKPT_SAVE_DIR, "final_model")) |
|
|
| if __name__ == "__main__": |
| main() |