File size: 1,763 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
from __future__ import annotations

import argparse

from .inference import generate_code, load_generator


def main() -> int:
    parser = argparse.ArgumentParser(description="Generate code with a Gemma LoRA adapter.")
    parser.add_argument("--base-model", default="google/gemma-3-1b-it", help="Base Hugging Face model id.")
    parser.add_argument("--adapter", default=None, help="Path to a trained LoRA adapter.")
    parser.add_argument("--instruction", required=True, help="Coding instruction.")
    parser.add_argument("--input", default="", help="Optional extra input/context.")
    parser.add_argument("--max-new-tokens", type=int, default=512)
    parser.add_argument("--temperature", type=float, default=0.2)
    parser.add_argument("--top-p", type=float, default=0.95)
    parser.add_argument("--quantization", choices=["none", "4bit", "8bit"], default="none")
    parser.add_argument("--dtype", choices=["auto", "float32", "float16", "bfloat16"], default="auto")
    parser.add_argument("--trust-remote-code", action="store_true")
    parser.add_argument("--disable-safety", action="store_true")
    args = parser.parse_args()

    model, tokenizer, torch = load_generator(
        args.base_model,
        adapter=args.adapter,
        quantization=args.quantization,
        dtype=args.dtype,
        trust_remote_code=args.trust_remote_code,
    )
    completion = generate_code(
        model,
        tokenizer,
        torch,
        instruction=args.instruction,
        input_text=args.input,
        max_new_tokens=args.max_new_tokens,
        temperature=args.temperature,
        top_p=args.top_p,
        safety=not args.disable_safety,
    )
    print(completion)
    return 0


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