| |
| """Run deterministic inference with the published t5-smaller checkpoint.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
|
|
| import torch |
| from transformers import AutoModelForSeq2SeqLM, AutoTokenizer |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("prompt", nargs="?", default="translate English to German: How old are you?") |
| parser.add_argument("--model", default="ShinpacheShimura/t5-smaller") |
| parser.add_argument("--subfolder", default="optimized-flan-t5-small") |
| parser.add_argument("--max-new-tokens", type=int, default=64) |
| args = parser.parse_args() |
|
|
| common = {"subfolder": args.subfolder} if args.subfolder else {} |
| tokenizer = AutoTokenizer.from_pretrained(args.model, **common) |
| model = AutoModelForSeq2SeqLM.from_pretrained(args.model, device_map="auto", **common) |
| inputs = tokenizer(args.prompt, return_tensors="pt").to(model.device) |
| with torch.inference_mode(): |
| output_ids = model.generate(**inputs, max_new_tokens=args.max_new_tokens, do_sample=False) |
| print(tokenizer.decode(output_ids[0], skip_special_tokens=True)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|