199 lines
6.0 KiB
Python
199 lines
6.0 KiB
Python
import argparse
|
|
import json
|
|
import re
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
|
|
from model import MiniLM, ModelConfig
|
|
|
|
|
|
def parse_args() -> argparse.Namespace:
|
|
p = argparse.ArgumentParser(description="Run inference with a trained mini-10m checkpoint")
|
|
p.add_argument("--checkpoint", type=str, default=None, help="Path to checkpoint .pt file")
|
|
p.add_argument("--run-dir", type=str, default="runs/tiny10m", help="Run dir containing step_*.pt")
|
|
p.add_argument("--prompt", type=str, default="Once upon a time")
|
|
p.add_argument("--max-new-tokens", type=int, default=200)
|
|
p.add_argument("--temperature", type=float, default=0.8)
|
|
p.add_argument("--top-k", type=int, default=40)
|
|
p.add_argument("--top-p", type=float, default=0.95)
|
|
p.add_argument("--seed", type=int, default=1337)
|
|
p.add_argument("--device", type=str, default="auto")
|
|
p.add_argument("--dtype", type=str, default="auto", choices=["auto", "float32", "bfloat16"])
|
|
p.add_argument(
|
|
"--max-seq-len",
|
|
type=int,
|
|
default=None,
|
|
help="Override model max context window. If omitted, tries run config then defaults to 16384.",
|
|
)
|
|
return p.parse_args()
|
|
|
|
|
|
def detect_device(user_value: str) -> torch.device:
|
|
if user_value != "auto":
|
|
return torch.device(user_value)
|
|
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
|
|
|
|
def detect_dtype(user_value: str, device: torch.device) -> torch.dtype:
|
|
if user_value == "float32":
|
|
return torch.float32
|
|
if user_value == "bfloat16":
|
|
return torch.bfloat16
|
|
if device.type == "cuda" and torch.cuda.is_bf16_supported():
|
|
return torch.bfloat16
|
|
return torch.float32
|
|
|
|
|
|
def encode_byte_level(text: str) -> list[int]:
|
|
return list(text.encode("utf-8", errors="ignore"))
|
|
|
|
|
|
def decode_byte_level(tokens: list[int]) -> str:
|
|
return bytes(tokens).decode("utf-8", errors="replace")
|
|
|
|
|
|
def find_latest_checkpoint(run_dir: Path) -> Path:
|
|
ckpts = list(run_dir.glob("step_*.pt"))
|
|
if not ckpts:
|
|
raise FileNotFoundError(f"No checkpoints found in {run_dir}")
|
|
|
|
def step_num(path: Path) -> int:
|
|
m = re.search(r"step_(\d+)\.pt$", path.name)
|
|
return int(m.group(1)) if m else -1
|
|
|
|
return max(ckpts, key=step_num)
|
|
|
|
|
|
def maybe_load_run_config(run_dir: Path) -> dict:
|
|
cfg_path = run_dir / "config.json"
|
|
if not cfg_path.exists():
|
|
return {}
|
|
with open(cfg_path, "r", encoding="utf-8") as f:
|
|
return json.load(f)
|
|
|
|
|
|
def build_model(max_seq_len: int) -> MiniLM:
|
|
cfg = ModelConfig(
|
|
vocab_size=256,
|
|
d_model=352,
|
|
n_heads=8,
|
|
n_layers=7,
|
|
ffn_mult=4,
|
|
max_seq_len=max_seq_len,
|
|
dropout=0.0,
|
|
)
|
|
return MiniLM(cfg)
|
|
|
|
|
|
def sample_next_token(
|
|
logits: torch.Tensor,
|
|
temperature: float,
|
|
top_k: int,
|
|
top_p: float,
|
|
) -> int:
|
|
if temperature <= 0:
|
|
return int(torch.argmax(logits).item())
|
|
|
|
logits = logits / temperature
|
|
|
|
if top_k > 0:
|
|
k = min(top_k, logits.size(-1))
|
|
values, _ = torch.topk(logits, k)
|
|
cutoff = values[..., -1]
|
|
logits = torch.where(logits < cutoff, torch.full_like(logits, float("-inf")), logits)
|
|
|
|
if 0.0 < top_p < 1.0:
|
|
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
|
|
sorted_probs = torch.softmax(sorted_logits, dim=-1)
|
|
cumulative_probs = torch.cumsum(sorted_probs, dim=-1)
|
|
|
|
sorted_mask = cumulative_probs > top_p
|
|
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
|
|
sorted_mask[..., 0] = False
|
|
|
|
mask = torch.zeros_like(sorted_mask, dtype=torch.bool)
|
|
mask.scatter_(0, sorted_indices, sorted_mask)
|
|
logits = torch.where(mask, torch.full_like(logits, float("-inf")), logits)
|
|
|
|
probs = torch.softmax(logits, dim=-1)
|
|
next_token = torch.multinomial(probs, num_samples=1)
|
|
return int(next_token.item())
|
|
|
|
|
|
@torch.no_grad()
|
|
def generate(
|
|
model: MiniLM,
|
|
prompt_ids: list[int],
|
|
max_new_tokens: int,
|
|
temperature: float,
|
|
top_k: int,
|
|
top_p: float,
|
|
device: torch.device,
|
|
amp_dtype: torch.dtype,
|
|
) -> list[int]:
|
|
out_ids = list(prompt_ids)
|
|
|
|
for _ in range(max_new_tokens):
|
|
context = out_ids[-model.cfg.max_seq_len :]
|
|
x = torch.tensor(context, dtype=torch.long, device=device).unsqueeze(0)
|
|
with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=device.type == "cuda"):
|
|
logits, _ = model(x)
|
|
next_logits = logits[0, -1]
|
|
next_id = sample_next_token(next_logits, temperature, top_k, top_p)
|
|
out_ids.append(next_id)
|
|
|
|
return out_ids
|
|
|
|
|
|
def main():
|
|
args = parse_args()
|
|
torch.manual_seed(args.seed)
|
|
torch.cuda.manual_seed_all(args.seed)
|
|
|
|
device = detect_device(args.device)
|
|
amp_dtype = detect_dtype(args.dtype, device)
|
|
|
|
run_dir = Path(args.run_dir)
|
|
ckpt_path = Path(args.checkpoint) if args.checkpoint else find_latest_checkpoint(run_dir)
|
|
run_cfg = maybe_load_run_config(ckpt_path.parent)
|
|
|
|
max_seq_len = args.max_seq_len
|
|
if max_seq_len is None:
|
|
max_seq_len = int(run_cfg.get("max_seq_len", 16_384))
|
|
|
|
model = build_model(max_seq_len=max_seq_len).to(device)
|
|
|
|
ckpt = torch.load(ckpt_path, map_location=device)
|
|
state_dict = ckpt["model"] if isinstance(ckpt, dict) and "model" in ckpt else ckpt
|
|
model.load_state_dict(state_dict, strict=True)
|
|
model.eval()
|
|
|
|
prompt_ids = encode_byte_level(args.prompt)
|
|
if not prompt_ids:
|
|
raise ValueError("Prompt produced no byte tokens. Provide a non-empty prompt.")
|
|
|
|
out_ids = generate(
|
|
model=model,
|
|
prompt_ids=prompt_ids,
|
|
max_new_tokens=args.max_new_tokens,
|
|
temperature=args.temperature,
|
|
top_k=args.top_k,
|
|
top_p=args.top_p,
|
|
device=device,
|
|
amp_dtype=amp_dtype,
|
|
)
|
|
|
|
generated_suffix = out_ids[len(prompt_ids) :]
|
|
print(f"Device: {device}")
|
|
print(f"Checkpoint: {ckpt_path}")
|
|
print(f"Max context: {max_seq_len}")
|
|
print("--- PROMPT ---")
|
|
print(args.prompt)
|
|
print("--- GENERATED ---")
|
|
print(decode_byte_level(generated_suffix))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|