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.json" # CKPT_SAVE_DIR = "/workspace/ckpt/pmnet" # CKPT_START_DIR = "/workspace/ckpt_start" # compile_model = True 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) # model size 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()