pmnet / copy_paste /train_copy_paste.py
phasorkinetics's picture
Upload 32 files
b24b632 verified
Raw
History Blame Contribute Delete
4.13 kB
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()