File size: 6,635 Bytes
c1de90b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
from __future__ import annotations

import argparse
import inspect
from pathlib import Path

from .data import TokenizedCodeDataset, load_jsonl_records, split_records
from .modeling import load_model_and_tokenizer


def _csv(value: str) -> list[str]:
    return [part.strip() for part in value.split(",") if part.strip()]


def _import_training_stack():
    try:
        from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
        from transformers import DataCollatorForSeq2Seq, Trainer, TrainingArguments
    except Exception as exc:
        raise RuntimeError(
            "Missing training dependencies. Install them with: pip install -r requirements.txt"
        ) from exc
    return LoraConfig, get_peft_model, prepare_model_for_kbit_training, DataCollatorForSeq2Seq, Trainer, TrainingArguments


def _build_training_args(TrainingArguments, args, torch):
    signature = inspect.signature(TrainingArguments.__init__).parameters
    cuda = torch.cuda.is_available()
    use_bf16 = cuda and torch.cuda.is_bf16_supported()
    use_fp16 = cuda and not use_bf16

    kwargs = {
        "output_dir": args.output_dir,
        "per_device_train_batch_size": args.batch_size,
        "gradient_accumulation_steps": args.gradient_accumulation_steps,
        "num_train_epochs": args.epochs,
        "learning_rate": args.learning_rate,
        "logging_steps": args.logging_steps,
        "save_steps": args.save_steps,
        "save_total_limit": args.save_total_limit,
        "warmup_ratio": args.warmup_ratio,
        "weight_decay": args.weight_decay,
        "optim": "adamw_torch",
        "report_to": "none",
        "remove_unused_columns": False,
        "gradient_checkpointing": args.gradient_checkpointing,
    }

    if "bf16" in signature:
        kwargs["bf16"] = use_bf16
    if "fp16" in signature:
        kwargs["fp16"] = use_fp16

    if args.validation_size > 0:
        strategy_name = "eval_strategy" if "eval_strategy" in signature else "evaluation_strategy"
        kwargs[strategy_name] = "steps"
        kwargs["eval_steps"] = args.eval_steps

    return TrainingArguments(**kwargs)


def main() -> int:
    parser = argparse.ArgumentParser(description="Fine-tune Gemma for code generation with LoRA.")
    parser.add_argument("--data", required=True, help="Path to JSONL training data.")
    parser.add_argument("--base-model", default="google/gemma-3-1b-it", help="Base Hugging Face model id.")
    parser.add_argument("--output-dir", default="outputs/gemma-code-lora", help="Where to save the LoRA adapter.")
    parser.add_argument("--max-records", type=int, default=None, help="Optional limit for quick tests.")
    parser.add_argument("--validation-size", type=float, default=0.0, help="Fraction of data for validation.")
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--max-length", type=int, default=2048)
    parser.add_argument("--epochs", type=float, default=2.0)
    parser.add_argument("--batch-size", type=int, default=1)
    parser.add_argument("--gradient-accumulation-steps", type=int, default=8)
    parser.add_argument("--learning-rate", type=float, default=2e-4)
    parser.add_argument("--warmup-ratio", type=float, default=0.03)
    parser.add_argument("--weight-decay", type=float, default=0.0)
    parser.add_argument("--logging-steps", type=int, default=10)
    parser.add_argument("--save-steps", type=int, default=100)
    parser.add_argument("--eval-steps", type=int, default=100)
    parser.add_argument("--save-total-limit", type=int, default=2)
    parser.add_argument("--lora-r", type=int, default=16)
    parser.add_argument("--lora-alpha", type=int, default=32)
    parser.add_argument("--lora-dropout", type=float, default=0.05)
    parser.add_argument(
        "--target-modules",
        default="q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj",
        help="Comma-separated LoRA target module names.",
    )
    parser.add_argument("--quantization", choices=["none", "4bit", "8bit"], default="none")
    parser.add_argument("--dtype", choices=["auto", "float32", "float16", "bfloat16"], default="auto")
    parser.add_argument("--gradient-checkpointing", action="store_true")
    parser.add_argument("--trust-remote-code", action="store_true")
    parser.add_argument("--resume-from-checkpoint", default=None)
    args = parser.parse_args()

    LoraConfig, get_peft_model, prepare_model_for_kbit_training, DataCollatorForSeq2Seq, Trainer, TrainingArguments = (
        _import_training_stack()
    )

    records = load_jsonl_records(args.data, max_records=args.max_records)
    train_records, eval_records = split_records(records, args.validation_size, args.seed)

    model, tokenizer, torch = load_model_and_tokenizer(
        args.base_model,
        quantization=args.quantization,
        dtype=args.dtype,
        trust_remote_code=args.trust_remote_code,
        for_training=True,
    )

    model.config.use_cache = False
    if args.quantization != "none":
        model = prepare_model_for_kbit_training(model)

    lora_config = LoraConfig(
        task_type="CAUSAL_LM",
        r=args.lora_r,
        lora_alpha=args.lora_alpha,
        lora_dropout=args.lora_dropout,
        target_modules=_csv(args.target_modules),
    )
    model = get_peft_model(model, lora_config)
    model.print_trainable_parameters()

    train_dataset = TokenizedCodeDataset(train_records, tokenizer, args.max_length)
    eval_dataset = TokenizedCodeDataset(eval_records, tokenizer, args.max_length) if eval_records else None
    data_collator = DataCollatorForSeq2Seq(
        tokenizer=tokenizer,
        model=model,
        label_pad_token_id=-100,
        pad_to_multiple_of=8,
    )

    training_args = _build_training_args(TrainingArguments, args, torch)
    trainer_kwargs = {
        "model": model,
        "args": training_args,
        "train_dataset": train_dataset,
        "eval_dataset": eval_dataset,
        "data_collator": data_collator,
    }

    trainer_signature = inspect.signature(Trainer.__init__).parameters
    if "processing_class" in trainer_signature:
        trainer_kwargs["processing_class"] = tokenizer
    else:
        trainer_kwargs["tokenizer"] = tokenizer

    trainer = Trainer(**trainer_kwargs)
    trainer.train(resume_from_checkpoint=args.resume_from_checkpoint)

    output_dir = Path(args.output_dir)
    output_dir.mkdir(parents=True, exist_ok=True)
    trainer.save_model(str(output_dir))
    tokenizer.save_pretrained(str(output_dir))
    print(f"Saved LoRA adapter to {output_dir}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())