71 lines
2.6 KiB
Python
71 lines
2.6 KiB
Python
import os, argparse
|
|
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
os.environ.setdefault("HF_HOME", os.path.join(HERE, "hf_cache"))
|
|
|
|
from datasets import load_dataset
|
|
from peft import LoraConfig
|
|
from trl import SFTConfig, SFTTrainer
|
|
|
|
MODEL = "Qwen/Qwen3-0.6B"
|
|
TRAIN = os.path.join(HERE, "data", "sft_train.jsonl")
|
|
VAL = os.path.join(HERE, "data", "sft_val.jsonl")
|
|
|
|
QWEN_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj",
|
|
"gate_proj", "up_proj", "down_proj"]
|
|
|
|
def main():
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--lr", type=float, default=2e-4)
|
|
ap.add_argument("--epochs", type=float, default=3)
|
|
ap.add_argument("--frac", type=float, default=1.0) # 0.1 for LR sweep
|
|
ap.add_argument("--dora", action="store_true") # run 1 = off
|
|
ap.add_argument("--out", default=os.path.join(HERE, "adapters", "sft-lora"))
|
|
ap.add_argument("--eval_steps", type=int, default=50)
|
|
ap.add_argument("--logging_steps", type=int, default=10)
|
|
ap.add_argument("--batch", type=int, default=8)
|
|
ap.add_argument("--grad_accum", type=int, default=4)
|
|
ap.add_argument("--no_grad_ckpt", action="store_true")
|
|
a = ap.parse_args()
|
|
|
|
train = load_dataset("json", data_files=TRAIN, split="train")
|
|
val = load_dataset("json", data_files=VAL, split="train")
|
|
if a.frac < 1.0:
|
|
train = train.select(range(int(len(train) * a.frac)))
|
|
|
|
peft = LoraConfig(r=16, lora_alpha=32, lora_dropout=0.05,
|
|
bias="none", task_type="CAUSAL_LM",
|
|
target_modules=QWEN_TARGETS, use_dora=a.dora)
|
|
|
|
cfg = SFTConfig(
|
|
output_dir=a.out,
|
|
model_init_kwargs={"dtype": "bfloat16"},
|
|
max_length=512,
|
|
packing=False,
|
|
assistant_only_loss=True,
|
|
use_liger_kernel=False,
|
|
per_device_train_batch_size=a.batch,
|
|
per_device_eval_batch_size=a.batch,
|
|
gradient_accumulation_steps=a.grad_accum,
|
|
num_train_epochs=a.epochs,
|
|
learning_rate=a.lr,
|
|
lr_scheduler_type="cosine",
|
|
warmup_ratio=0.03,
|
|
bf16=True,
|
|
gradient_checkpointing=not a.no_grad_ckpt,
|
|
logging_steps=a.logging_steps,
|
|
eval_strategy="steps",
|
|
eval_steps=a.eval_steps,
|
|
save_strategy="epoch",
|
|
report_to="none",
|
|
)
|
|
|
|
trainer = SFTTrainer(model=MODEL, args=cfg,
|
|
train_dataset=train, eval_dataset=val,
|
|
peft_config=peft)
|
|
trainer.train()
|
|
trainer.save_model(a.out)
|
|
print("saved adapter ->", a.out)
|
|
print("final metrics:", trainer.state.log_history[-1])
|
|
|
|
if __name__ == "__main__":
|
|
main() |