Pluto-Lite-v1 / load_adapter.py
SanaeNya's picture
Upload folder using huggingface_hub
6cff49b verified
Raw History Blame Contribute Delete
3.49 kB
"""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()