Views
No views yet
1import os
2from unsloth import FastModel
3import torch
4from trl import SFTConfig, SFTTrainer
5from teich import mask_data, prepare_data
6
7MAX_SEQ_LEN = 32768
8MODEL_NAME = "Qwen/Qwen3.5-9B"
9OUTPUT_DIR = "/content/drive/MyDrive/Colab/outputs-qwen-tool-sft"
10HUB_REPO_ID = "armand0e/Qwen3.5-9B-Coder"
11HF_TOKEN = os.environ.get("HF_TOKEN", "")
12CHAT_TEMPLATE_PATH = "qwen3.5-chat-template.jinja"
13
14model, tokenizer = FastModel.from_pretrained(
15 model_name=MODEL_NAME,
16 max_seq_length=MAX_SEQ_LEN,
17 load_in_4bit=False,
18 load_in_8bit=False,
19 full_finetuning=False,
20 token=HF_TOKEN,
21)
22
23if CHAT_TEMPLATE_PATH:
24 with open(CHAT_TEMPLATE_PATH, "r", encoding="utf-8") as f:
25 custom_chat_template = f.read()
26 tokenizer.chat_template = custom_chat_template
27 if hasattr(tokenizer, "tokenizer") and tokenizer.tokenizer is not None:
28 tokenizer.tokenizer.chat_template = custom_chat_template
29
30model = FastModel.get_peft_model(
31 model,
32 finetune_vision_layers = False, # Turn off for just text!
33 finetune_language_layers = True, # Should leave on!
34 finetune_attention_modules = True, # Attention good for GRPO
35 finetune_mlp_modules = True, # Should leave on always!
36
37 r = 32, # Larger = higher accuracy, but might overfit
38 lora_alpha = 32, # Recommended alpha == r at least
39 lora_dropout = 0,
40 bias = "none",
41 random_state = 3407,
42)
43
44train_dataset = prepare_data(
45 {
46 "qwen3.7-max": {
47 "source": "armand0e/qwen3.7-max", # stupid typo i made and now this model wasn't trained on the qwen3.7-max traces :(
48 },
49 "chat": {
50 "source": "TeichAI/claude-4.5-opus-high-reasoning-250x",
51 },
52 "opus-pi-agent": {
53 "source": "armand0e/badlogicgames-pi-mono-opus-filtered",
54 },
55 "kimi-k2.6-claude-code": {
56 "source": "armand0e/kimi-k2.6-claude-code-traces",
57 },
58 "chat-2": {
59 "source": "TeichAI/Claude-Opus-4.6-Reasoning-887x"
60 },
61 "minimax-m3-claude-code": {
62 "source": "armand0e/minimax-m3-claude-code-traces"
63 },
64 "more-opus": {
65 "source": "armand0e/claude-opus-4.8-pi-traces"
66 }
67 },
68 tokenizer,
69 split="train",
70 hf_token=HF_TOKEN,
71 chat_template_kwargs={"enable_thinking": False, "preserve_thinking": True},
72 max_length=MAX_SEQ_LEN,
73 oversized_policy="trim_followups",
74 tokenize=True,
75 strict=True,
76)
77
78trainer = SFTTrainer(
79 model=model,
80 tokenizer=tokenizer,
81 train_dataset=train_dataset,
82 eval_dataset=None,
83 args=SFTConfig(
84 dataset_text_field="text",
85 dataset_num_proc=1,
86 max_length=MAX_SEQ_LEN,
87 packing=False,
88 per_device_train_batch_size=1,
89 gradient_accumulation_steps=8,
90 warmup_steps= 5,
91 num_train_epochs=1,
92 learning_rate=2e-4,
93 logging_steps=1,
94 save_strategy="epoch",
95 save_total_limit=3,
96 optim="adamw_8bit",
97 weight_decay=0.01,
98 #max_grad_norm=0.3,
99 lr_scheduler_type="linear",
100 output_dir=OUTPUT_DIR,
101 seed=3407,
102 report_to="none",
103 ),
104)
105
106trainer = mask_data(
107 trainer,
108 tokenizer=tokenizer,
109 train_on_reasoning=False,
110 train_on_final_answers=True,
111 train_on_tools=True,
112)
113
114print(trainer.train_dataset.preview())
115
116trainer_stats = trainer.train(resume_from_checkpoint=False)
117
118model.push_to_hub(f"{HUB_REPO_ID}-LoRA", token=HF_TOKEN)
119tokenizer.push_to_hub(f"{HUB_REPO_ID}-LoRA", token=HF_TOKEN)
120
121model.push_to_hub_merged(HUB_REPO_ID, tokenizer, save_method="merged_16bit", token=HF_TOKEN)