初始化项目,由ModelHub XC社区提供模型
Model: karthik-2905/AL1-model-B Source: Original Platform
This commit is contained in:
71
mlb/train_sft.py
Normal file
71
mlb/train_sft.py
Normal file
@@ -0,0 +1,71 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user