Views
No views yet
1def chatml_format(example):
2 # Initialize formatted system message
3 system = ""
4
5 # Check if 'system' field exists and is not None
6 if example.get('system'):
7 message = {"role": "system", "content": example['system']}
8 system = tokenizer.apply_chat_template([message], tokenize=False)
9
10 # Format instruction
11 message = {"role": "user", "content": example['prompt']}
12 prompt = tokenizer.apply_chat_template([message], tokenize=False, add_generation_prompt=True)
13
14 # Format chosen answer
15 chosen = example['chosen'] + "<|im_end|>\n"
16
17 # Format rejected answer
18 rejected = example['rejected'] + "<|im_end|>\n"
19
20 return {
21 "prompt": system + prompt,
22 "chosen": chosen,
23 "rejected": rejected,
24 }
25
26# Array of datasets to concat
27ds = [
28 "jondurbin/truthy-dpo-v0.1",
29 "ResplendentAI/NSFW_RP_Format_DPO",
30 "jondurbin/gutenberg-dpo-v0.1",
31 "flammenai/Date-DPO-v1"
32]
33
34# load_dataset and combine all
35loaded_datasets = [load_dataset(dataset_name, split='train') for dataset_name in ds]
36dataset = concatenate_datasets(loaded_datasets)
37
38# Save columns
39original_columns = dataset.column_names
40
41# Tokenizer
42tokenizer = AutoTokenizer.from_pretrained(model_name)
43tokenizer.pad_token = tokenizer.eos_token
44tokenizer.padding_side = "left"
45
46# Format dataset
47dataset = dataset.map(
48 chatml_format,
49 remove_columns=original_columns
50)1# LoRA configuration
2peft_config = LoraConfig(
3 r=16,
4 lora_alpha=16,
5 lora_dropout=0.05,
6 bias="none",
7 task_type="CAUSAL_LM",
8 target_modules=['k_proj', 'gate_proj', 'v_proj', 'up_proj', 'q_proj', 'o_proj', 'down_proj']
9)
10# Model to fine-tune
11model = AutoModelForCausalLM.from_pretrained(
12 model_name,
13 torch_dtype=torch.bfloat16,
14 load_in_4bit=True
15)
16model.config.use_cache = False
17# Reference model
18ref_model = AutoModelForCausalLM.from_pretrained(
19 model_name,
20 torch_dtype=torch.bfloat16,
21 load_in_4bit=True
22)
23# Training arguments
24training_args = TrainingArguments(
25 per_device_train_batch_size=2,
26 gradient_accumulation_steps=8,
27 gradient_checkpointing=True,
28 learning_rate=5e-5,
29 lr_scheduler_type="cosine",
30 max_steps=420,
31 save_strategy="no",
32 logging_steps=1,
33 output_dir=new_model,
34 optim="paged_adamw_32bit",
35 warmup_steps=100,
36 bf16=True,
37 report_to="wandb",
38)
39# Create DPO trainer
40dpo_trainer = DPOTrainer(
41 model,
42 ref_model,
43 args=training_args,
44 train_dataset=dataset,
45 tokenizer=tokenizer,
46 peft_config=peft_config,
47 beta=0.1,
48 max_prompt_length=2048,
49 max_length=4096,
50 force_use_ref_model=True
51)
52# Fine-tune model with DPO
53dpo_trainer.train()