1from datasets import Dataset
2from trl import GKDConfig, GKDTrainer
3from transformers import (
4 AutoModelForCausalLM,
5 AutoTokenizer,
6)
7from datasets import load_dataset
8from peft import LoraConfig
9
10parser = argparse.ArgumentParser()
11parser.add_argument("--temperature", type=float, default = 0.9)
12parser.add_argument("--lmbda", type=float, default = 0.5)
13parser.add_argument("--beta", type=float, default = 0.5)
14parser.add_argument("--max_new_tokens", type=int, default = 4096)
15parser.add_argument("--output_dir", type=str, default="gkd-model")
16parser.add_argument("--per_device_train_batch_size", type=int, default=1)
17parser.add_argument("--gradient_accumulation_steps", type=int, default=16)
18parser.add_argument("--gradient_checkpointing", action="store_true", default=False)
19parser.add_argument("--resume_from_checkpoint", action="store_true", default=False)
20parser.add_argument("--lora", action="store_true")
21args = parser.parse_args()
22
23qwq_dataset = load_dataset("amphora/QwQ-LongCoT-130K", split = "train")
24messages = []
25for each in qwq_dataset:
26 msg = [
27 {"role": "system", "content": "You are a helpful and harmless assistant. You are Qwen developed by Alibaba. You should think step-by-step."},
28 {"role": "user", "content": each["problem"]},
29 {"role": "assistant", "content": each["qwq"]},
30 ]
31 messages.append(msg)
32
33TRAIN_SPLIT_RATIO = 0.9
34train_size = int(TRAIN_SPLIT_RATIO * len(messages))
35eval_size = len(messages) - train_size
36
37tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2-0.5B-Instruct")
38
39
40
41
42# The teacher model to calculate the KL divergence against
43teacher_model = AutoModelForCausalLM.from_pretrained("Qwen/QwQ-32B-Preview", torch_dtype=torch.bfloat16, device_map="auto")
44teacher_model.lm_head.weight.data = teacher_model.lm_head.weight.data[:151936, :]
45teacher_model.lm_head.out_features = 151936
46
47
48
49# The model to optimise
50model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2-0.5B-Instruct", torch_dtype=torch.bfloat16, device_map="auto")
51
52
53
54### Real Dataset
55train_dataset = Dataset.from_dict({"messages":messages[:train_size]})
56eval_dataset = Dataset.from_dict({"messages":messages[train_size:]})
57training_args = GKDConfig(
58 output_dir=args.output_dir,
59 temperature=args.temperature,
60 lmbda=args.lmbda,
61 beta=args.beta,
62 max_new_tokens=args.max_new_tokens,
63 per_device_train_batch_size=args.per_device_train_batch_size,
64 gradient_accumulation_steps=args.gradient_accumulation_steps,
65 gradient_checkpointing = args.gradient_checkpointing,
66 save_steps = 100,
67 save_total_limit = 5
68 )
69
70lora_config = LoraConfig(
71 r=16,
72 lora_alpha=32,
73 lora_dropout=0.05,
74 bias="none",
75 task_type="CAUSAL_LM",
76)
77
78trainer = GKDTrainer(
79 model=model,
80 teacher_model=teacher_model,
81 args=training_args,
82 processing_class=tokenizer,
83 train_dataset=train_dataset,
84 eval_dataset=eval_dataset,
85 peft_config=lora_config if args.lora else None
86)
87trainer.train(resume_from_checkpoint=args.resume_from_checkpoint)