File size: 3,493 Bytes
6cff49b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Load the PlutoLite-v1 text adapter with its original completion format."""

import argparse
from pathlib import Path


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--adapter", default=str(Path(__file__).resolve().parent))
    parser.add_argument("--base-model", default="google/gemma-4-E2B")
    parser.add_argument("--local-files-only", action="store_true")
    parser.add_argument("--prompt", required=True)
    parser.add_argument("--max-new-tokens", type=int, default=256)
    args = parser.parse_args()
    if args.max_new_tokens < 1:
        parser.error("--max-new-tokens must be positive")

    import torch
    from accelerate.hooks import remove_hook_from_module
    from peft import PeftModel
    from transformers import AutoTokenizer, BitsAndBytesConfig, Gemma4ForCausalLM

    if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
        raise RuntimeError("This NF4 example requires a BF16-capable NVIDIA GPU.")
    revision = "d29ff6b45f081a49ee2733a859c9c9c2d95d1a6f"
    source_kwargs = dict(
        revision=revision,
        local_files_only=args.local_files_only,
        trust_remote_code=False,
    )
    tokenizer = AutoTokenizer.from_pretrained(args.base_model, **source_kwargs)
    if tokenizer.pad_token_id is None:
        tokenizer.pad_token = tokenizer.eos_token
    base, loading = Gemma4ForCausalLM.from_pretrained(
        args.base_model,
        **source_kwargs,
        dtype=torch.bfloat16,
        device_map={"": "cuda:0", "model.embed_tokens_per_layer": "cpu"},
        attn_implementation="sdpa",
        quantization_config=BitsAndBytesConfig(
            load_in_4bit=True,
            bnb_4bit_quant_type="nf4",
            bnb_4bit_use_double_quant=True,
            bnb_4bit_compute_dtype=torch.bfloat16,
            llm_int8_enable_fp32_cpu_offload=True,
        ),
        key_mapping={r"^model\.language_model\.": "model."},
        output_loading_info=True,
    )
    if any(loading.get(key) for key in ("missing_keys", "mismatched_keys", "error_msgs")):
        raise RuntimeError(f"Incomplete base text model loading: {loading}")

    # Keep the large PLE table on CPU and transfer only lookup results to CUDA.
    embedding = base.model.embed_tokens_per_layer
    remove_hook_from_module(embedding, recurse=True)
    embedding.to(device="cpu", dtype=torch.bfloat16)
    embedding.register_forward_pre_hook(lambda module, inputs: (inputs[0].to("cpu"),))
    embedding.register_forward_hook(lambda module, inputs, output: output.to("cuda:0"))
    placement = getattr(base, "hf_device_map", None)
    if placement is not None:
        del base.hf_device_map
    try:
        model = PeftModel.from_pretrained(
            base, args.adapter, is_trainable=False, local_files_only=args.local_files_only
        )
    finally:
        if placement is not None:
            base.hf_device_map = placement
    model.eval()
    prefix = (tokenizer.bos_token or "") + args.prompt
    inputs = tokenizer(prefix, add_special_tokens=False, return_tensors="pt").to("cuda:0")
    with torch.inference_mode():
        output = model.generate(
            **inputs,
            max_new_tokens=args.max_new_tokens,
            do_sample=False,
            pad_token_id=tokenizer.pad_token_id,
            eos_token_id=tokenizer.eos_token_id,
        )
    print(tokenizer.decode(output[0, inputs.input_ids.shape[1]:], skip_special_tokens=True))


if __name__ == "__main__":
    main()