Views
No views yet
1import os
2import argparse
3import torch
4from datasets import Dataset
5from trl import SFTConfig, SFTTrainer, DataCollatorForCompletionOnlyLM
6from transformers import (
7 AutoModelForCausalLM,
8 AutoTokenizer,
9)
10from datasets import load_dataset
11from peft import LoraConfig
12
13parser = argparse.ArgumentParser()
14parser.add_argument("--max_length", 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-2", 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.5-0.5B-Instruct")
38
39# The model to optimise
40model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-0.5B-Instruct", torch_dtype=torch.bfloat16, device_map="auto")
41
42
43
44### Real Dataset
45train_dataset = Dataset.from_dict({"messages":messages[:train_size]})
46eval_dataset = Dataset.from_dict({"messages":messages[train_size:]})
47training_args = SFTConfig(
48 output_dir=args.output_dir,
49 max_seq_length=args.max_length,
50 per_device_train_batch_size=args.per_device_train_batch_size,
51 gradient_accumulation_steps=args.gradient_accumulation_steps,
52 gradient_checkpointing = args.gradient_checkpointing,
53 save_steps = 100,
54 save_total_limit = 5
55 )
56
57lora_config = LoraConfig(
58 r=16,
59 lora_alpha=32,
60 lora_dropout=0.05,
61 bias="none",
62 task_type="CAUSAL_LM",
63)
64
65response_template = "<|im_start|>assistant\n"
66
67collator = DataCollatorForCompletionOnlyLM(response_template, tokenizer=tokenizer)
68
69trainer = SFTTrainer(
70 model=model,
71 args=training_args,
72 processing_class=tokenizer,
73 train_dataset=train_dataset,
74 eval_dataset=eval_dataset,
75 peft_config=lora_config if args.lora else None,
76 data_collator=collator,
77)
78trainer.train(resume_from_checkpoint=args.resume_from_checkpoint)amphora/QwQ-LongCoT-130K1import torch
2from transformers import AutoModelForCausalLM, AutoTokenizer
3# Model name
4model_name = "kz919/QwQ-0.5B-Distilled-SFT"
5# Load the model
6print(f"Starting to load the model {model_name} into memory")
7model = AutoModelForCausalLM.from_pretrained(
8 model_name,
9 torch_dtype=torch.bfloat16,
10 device_map={"": 0}
11)
12# Load the tokenizer
13tokenizer = AutoTokenizer.from_pretrained(model_name)
14# Define the prompt
15prompt = "How many r in strawberry."
16messages = [
17 {"role": "system", "content": "You are a helpful and harmless assistant. You are Qwen developed by Alibaba. You should think step-by-step."},
18 {"role": "user", "content": prompt}
19]
20# Tokenize the input
21text = tokenizer.apply_chat_template(
22 messages,
23 tokenize=False,
24 add_generation_prompt=True
25)
26model_inputs = tokenizer([text], return_tensors="pt").to(model.device)
27# Generate a response
28generated_ids = model.generate(
29 **model_inputs,
30 max_new_tokens=4096
31)
32generated_ids = [
33 output_ids[len(input_ids):] for input_ids, output_ids in zip(model_inputs.input_ids, generated_ids)
34]
35# Decode the response
36response = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
37print(response)