tiny-ced / generate.py
Ne30Charm's picture
Upload folder using huggingface_hub
d5b9248 verified
Raw
History Blame Contribute Delete
2.14 kB
import argparse
import json
from pathlib import Path
import torch
from safetensors.torch import load_model
from tokenizers import Tokenizer
from model import CED
def main():
parser = argparse.ArgumentParser()
parser.add_argument(
"--model-dir",
type=Path,
default=Path(__file__).resolve().parent,
)
parser.add_argument("--text", default="Once upon a time, a little rabbit")
parser.add_argument("--device", choices=["cpu", "cuda"], default="cpu")
parser.add_argument("--temperature", type=float, default=0.8)
parser.add_argument("--max-new-tokens", type=int, default=150)
args = parser.parse_args()
config = json.loads((args.model_dir / "config.json").read_text())
model = CED(
vocab=config["vocab_size"],
dim=config["hidden_size"],
heads=config["num_attention_heads"],
ff=config["intermediate_size"],
layers=config["num_hidden_layers"],
window=config["local_window_size"],
)
load_model(model, args.model_dir / "model.safetensors")
model = model.to(args.device).eval()
tokenizer = Tokenizer.from_file(str(args.model_dir / "tokenizer.json"))
ids = [config["bos_token_id"]] + tokenizer.encode(args.text).ids
context_length = config["max_position_embeddings"]
if len(ids) > context_length:
raise ValueError(f"Prompt exceeds the {context_length}-token context length")
torch.manual_seed(123)
with torch.inference_mode():
for _ in range(min(args.max_new_tokens, context_length - len(ids))):
logits = model(torch.tensor([ids], device=args.device))[0, -1].float()
logits[
[config["pad_token_id"], config["bos_token_id"], config["unk_token_id"]]
] = -torch.inf
if args.temperature <= 0:
token = int(logits.argmax())
else:
token = int(torch.multinomial((logits / args.temperature).softmax(-1), 1))
ids.append(token)
if token == config["eos_token_id"]:
break
print(tokenizer.decode(ids))
if __name__ == "__main__":
main()