File size: 3,556 Bytes
66ccdde | 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 | from __future__ import annotations
import argparse
import hashlib
from pathlib import Path
import torch
from safetensors.torch import load_model
from tokenizers import Tokenizer
from barunlm import BarunConfig, BarunLM
ROOT = Path(__file__).resolve().parent
EXPECTED_SHA256 = {
"model.safetensors": "f2a7c88b9f2c2e3584809081407ab136795d82e30e89b730e007781c45d01447",
"barun_config.json": "9b3a1d71baa95a198744d250f9629231738d942570b8685c44307fd83dd33565",
"tokenizer.json": "70ded9605fccd09c2340ca7e225361eab0ae8b4dbbb0d6e26343ab5183979db6",
}
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Generate text with BarunLM-35M.")
parser.add_argument("--prompt", default="The future of small language models is")
parser.add_argument("--max-new-tokens", type=int, default=48)
parser.add_argument("--temperature", type=float, default=0.8)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--verify-only", action="store_true")
return parser.parse_args()
def main() -> None:
args = parse_args()
for name, expected in EXPECTED_SHA256.items():
actual = sha256(ROOT / name)
if actual != expected:
raise RuntimeError(f"{name} SHA-256 mismatch: expected {expected}, got {actual}")
config = BarunConfig.from_json(ROOT / "barun_config.json")
model = BarunLM(config)
missing, unexpected = load_model(model, ROOT / "model.safetensors", strict=False)
if missing or unexpected:
raise RuntimeError(f"checkpoint mismatch: missing={missing}, unexpected={unexpected}")
count = model.parameter_counts()["total"]
if count != 35_072_768:
raise RuntimeError(f"unexpected parameter count: {count}")
print(f"verified parameters={count} checkpoint_sha256={EXPECTED_SHA256['model.safetensors']}")
if args.verify_only:
return
if args.max_new_tokens < 1:
raise ValueError("--max-new-tokens must be positive")
if args.temperature < 0:
raise ValueError("--temperature cannot be negative")
device_name = "cuda" if args.device == "auto" and torch.cuda.is_available() else args.device
if device_name == "auto":
device_name = "cpu"
if device_name == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is unavailable")
device = torch.device(device_name)
dtype = torch.bfloat16 if device.type == "cuda" else torch.float32
model.to(device=device, dtype=dtype).eval()
tokenizer = Tokenizer.from_file(str(ROOT / "tokenizer.json"))
prompt_ids = tokenizer.encode(args.prompt, add_special_tokens=False).ids
if not prompt_ids:
raise ValueError("prompt must encode to at least one token")
if len(prompt_ids) + args.max_new_tokens > config.max_seq_len:
raise ValueError("prompt and continuation exceed the 2,048-token context")
input_ids = torch.tensor([prompt_ids], dtype=torch.long, device=device)
with torch.inference_mode():
output_ids = model.generate(
input_ids,
max_new_tokens=args.max_new_tokens,
temperature=args.temperature,
)
print(tokenizer.decode(output_ids[0].tolist(), skip_special_tokens=True))
if __name__ == "__main__":
main()
|