File size: 3,186 Bytes
4062845
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
from pathlib import Path
import torch
from datasets import load_dataset
from peft import LoraConfig, TaskType, get_peft_model
from torch.nn.utils.rnn import pad_sequence
from transformers import AutoConfig, AutoModelForSeq2SeqLM, AutoTokenizer, Seq2SeqTrainer, Seq2SeqTrainingArguments, set_seed


def encode(row):
    prompt = row["instruction"] + ("\n\n" + row["input"] if row["input"] else "")
    source = tokenizer(prompt, add_special_tokens=False).input_ids[:510]
    target = tokenizer(row["output"], add_special_tokens=False).input_ids[:126]
    return {"input_ids": [mode_id, *source, span_id],
            "decoder_input_ids": [config.decoder.bos_token_id, span_id, *target],
            "labels": [-100, *target, tokenizer.eos_token_id]}


def collate(rows):
    batch = {key: pad_sequence([torch.tensor(row[key]) for row in rows], batch_first=True, padding_value=fill)
             for key, fill in (("input_ids", tokenizer.pad_token_id), ("decoder_input_ids", tokenizer.pad_token_id), ("labels", -100))}
    for key, mask in (("input_ids", "attention_mask"), ("decoder_input_ids", "decoder_attention_mask")):
        batch[mask] = torch.arange(batch[key].shape[1])[None, :] < torch.tensor([len(r[key]) for r in rows])[:, None]
    return batch


if __name__ == "__main__":
    model_path = "/path/to/model"
    output_dir = Path("/path/to/output-adapter")
    max_steps = 20
    attention = "eager"
    assert torch.cuda.is_available()
    torch.set_num_threads(2)
    set_seed(42)
    tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
    tokenizer.pad_token = tokenizer.eos_token
    tokenizer.padding_side = "right"
    config = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
    config.use_cache = config.decoder.use_cache = False
    mode_id, span_id = tokenizer.convert_tokens_to_ids(["[_S_]", "<SPAN#0>"])
    data = load_dataset("tatsu-lab/alpaca", revision="dce01c9b08f87459cf36a430d809084718273017", split="train")
    data = data.map(encode, remove_columns=data.column_names)
    model = AutoModelForSeq2SeqLM.from_pretrained(
        model_path, config=config, trust_remote_code=True, dtype=torch.bfloat16,
        attn_implementation=attention, device_map={"": "cuda:0"},
    )
    model = get_peft_model(model, LoraConfig(
        task_type=TaskType.SEQ_2_SEQ_LM, r=16, lora_alpha=32, lora_dropout=0.05,
        target_modules=["q_proj", "k_proj", "v_proj", "o_proj"], bias="none",
    ))
    trainer = Seq2SeqTrainer(
        model=model, train_dataset=data, data_collator=collate, processing_class=tokenizer,
        args=Seq2SeqTrainingArguments(
            output_dir=str(output_dir), num_train_epochs=1, max_steps=max_steps,
            bf16=True, per_device_train_batch_size=1, gradient_accumulation_steps=4,
            learning_rate=2e-4, optim="adamw_torch", lr_scheduler_type="constant",
            gradient_checkpointing=True, gradient_checkpointing_kwargs={"use_reentrant": False},
            save_strategy="no", logging_steps=1, report_to="none", remove_unused_columns=False,
            dataloader_num_workers=0, seed=42, data_seed=42,
        ),
    )
    trainer.train()
    trainer.save_model()