Download example_inference.py from ShinpacheShimura/t5-smaller: direct link, hf CLI and curl.
- Browser
- Download file 1.2 kB
-
https://huggingface.co/ShinpacheShimura/t5-smaller/resolve/main/example_inference.py
- Command line
-
hf download hf://ShinpacheShimura/t5-smaller/example_inference.py
-
curl -L -o example_inference.py https://huggingface.co/ShinpacheShimura/t5-smaller/resolve/main/example_inference.py
1.2 kB
| #!/usr/bin/env python3 | |
| """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() | |