File size: 5,813 Bytes
efdf5ed | 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 164 165 166 167 168 169 170 171 172 | #!/usr/bin/env python3
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 Hamid Wakili <hamid@ideployed.com>
"""Run offline generation from the released Safetensors checkpoint."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Sequence
import torch
import torch.nn.functional as F
from safetensors.torch import load_file
from bpe_tokenizer import BPE_Tokenizer
from german_transformer_model import GermanGPT, GermanGPTConfig
def load_runtime(
model_dir: str | Path,
device_name: str = "auto",
) -> tuple[BPE_Tokenizer, GermanGPT, torch.device]:
"""Load the tokenizer, configuration, and Safetensors weights."""
root = Path(model_dir)
device = torch.device(
"cuda" if device_name == "auto" and torch.cuda.is_available()
else "cpu" if device_name == "auto"
else device_name
)
with (root / "config.json").open("r", encoding="utf-8") as file:
config = GermanGPTConfig.from_dict(json.load(file))
tokenizer = BPE_Tokenizer.load(root / "tokenizer" / "de_bpe_32k")
if tokenizer.vocabulary_size() != config.vocab_size:
raise ValueError(
f"Tokenizer size {tokenizer.vocabulary_size()} does not match "
f"model vocab_size {config.vocab_size}"
)
model = GermanGPT(config)
state_dict = load_file(root / "model.safetensors", device="cpu")
model.load_state_dict(state_dict, strict=True)
model.to(device).eval()
return tokenizer, model, device
def filter_logits(logits: torch.Tensor, top_k: int, top_p: float) -> torch.Tensor:
"""Apply top-k and nucleus filtering to one vocabulary-sized logit vector."""
filtered = logits.clone()
if top_k > 0:
cutoff = torch.topk(filtered, min(top_k, filtered.numel())).values[-1]
filtered[filtered < cutoff] = float("-inf")
if top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(filtered, descending=True)
probabilities = F.softmax(sorted_logits, dim=-1)
cumulative = torch.cumsum(probabilities, dim=-1)
remove = cumulative > top_p
remove[1:] = remove[:-1].clone()
remove[0] = False
filtered[sorted_indices[remove]] = float("-inf")
return filtered
def generate(
prompt: str,
tokenizer: BPE_Tokenizer,
model: GermanGPT,
device: torch.device,
max_new_tokens: int,
temperature: float,
top_k: int,
top_p: float,
repetition_penalty: float,
) -> str:
"""Generate a completion with greedy or temperature-based sampling."""
token_ids = tokenizer.encode(prompt)
if not token_ids:
token_ids = [tokenizer.special_tokens.get("<s>", 2)]
generated = list(token_ids)
eos_id = tokenizer.special_tokens.get("</s>", 3)
with torch.inference_mode():
for _ in range(max_new_tokens):
context = generated[-model.config.context_len :]
input_ids = torch.tensor([context], dtype=torch.long, device=device)
logits, _ = model(input_ids)
next_logits = logits[0, -1].float()
if repetition_penalty != 1.0:
for token_id in set(generated):
value = next_logits[token_id]
next_logits[token_id] = (
value * repetition_penalty
if value < 0
else value / repetition_penalty
)
if temperature <= 0:
next_id = int(torch.argmax(next_logits).item())
else:
next_logits = filter_logits(
next_logits / temperature,
top_k=top_k,
top_p=top_p,
)
probabilities = F.softmax(next_logits, dim=-1)
next_id = int(torch.multinomial(probabilities, 1).item())
generated.append(next_id)
if next_id == eos_id:
break
return tokenizer.decode(generated)
def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
"""Parse command-line arguments."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-dir", default=".")
parser.add_argument("--prompt", required=True)
parser.add_argument("--device", default="auto")
parser.add_argument("--max-new-tokens", type=int, default=80)
parser.add_argument("--temperature", type=float, default=0.8)
parser.add_argument("--top-k", type=int, default=40)
parser.add_argument("--top-p", type=float, default=0.9)
parser.add_argument("--repetition-penalty", type=float, default=1.1)
return parser.parse_args(argv)
def main(argv: Sequence[str] | None = None) -> int:
"""Load the release and print one generated completion."""
args = parse_args(argv)
if args.max_new_tokens < 0:
raise ValueError("--max-new-tokens must be non-negative")
if args.temperature < 0:
raise ValueError("--temperature must be non-negative")
if args.top_k < 0:
raise ValueError("--top-k must be non-negative")
if not 0 < args.top_p <= 1:
raise ValueError("--top-p must be in (0, 1]")
if args.repetition_penalty <= 0:
raise ValueError("--repetition-penalty must be positive")
tokenizer, model, device = load_runtime(args.model_dir, args.device)
print(
generate(
prompt=args.prompt,
tokenizer=tokenizer,
model=model,
device=device,
max_new_tokens=args.max_new_tokens,
temperature=args.temperature,
top_k=args.top_k,
top_p=args.top_p,
repetition_penalty=args.repetition_penalty,
)
)
return 0
if __name__ == "__main__":
raise SystemExit(main())
|