"""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()