Download code/generate.py from ukung/semantic-lite-2-decoder-smoke-test: direct link, hf CLI and curl.
- Browser
- Download file 1.96 kB
-
https://huggingface.co/ukung/semantic-lite-2-decoder-smoke-test/resolve/main/code/generate.py
- Command line
-
hf download hf://ukung/semantic-lite-2-decoder-smoke-test/code/generate.py
-
curl -L -o generate.py https://huggingface.co/ukung/semantic-lite-2-decoder-smoke-test/resolve/main/code/generate.py
1.96 kB
| """ | |
| Inference: load the trained decoder and generate from a prompt. | |
| Run: python generate.py "your problem text here" | |
| python generate.py # runs the built-in examples | |
| """ | |
| import sys | |
| import torch | |
| from data import MAX_SRC_LEN | |
| from encoder_loader import load_encoder | |
| from model import SemanticConditionedDecoder | |
| DECODER_WEIGHTS = "decoder_weights.pt" | |
| def load_model(device="cuda", weights=DECODER_WEIGHTS): | |
| encoder, tokenizer = load_encoder(device=device) | |
| model = SemanticConditionedDecoder(encoder=encoder, tokenizer=tokenizer).to(device) | |
| state = torch.load(weights, map_location=device) | |
| model.load_state_dict(state, strict=False) | |
| model.eval() | |
| return model, tokenizer | |
| def generate_text(model, tokenizer, prompt, device="cuda", | |
| max_new_tokens=256, temperature=0.8): | |
| enc = tokenizer(prompt, padding=True, truncation=True, | |
| max_length=MAX_SRC_LEN, return_tensors="pt") | |
| input_ids = enc["input_ids"].to(device) | |
| attention_mask = enc["attention_mask"].to(device) | |
| out = model.generate(input_ids, attention_mask, | |
| max_new_tokens=max_new_tokens, temperature=temperature) | |
| return tokenizer.decode(out[0][1:], skip_special_tokens=True) # drop BOS | |
| EXAMPLES = [ | |
| "252 fifth-grade students and 8 teachers are going on a field trip. " | |
| "If renting a 41-seater bus costs 300,000 won and the highway toll per bus " | |
| "is 7,500 won, how much does it cost to rent the buses and pay the tolls?", | |
| ] | |
| def main(): | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model, tokenizer = load_model(device=device) | |
| prompts = sys.argv[1:] or EXAMPLES | |
| for i, prompt in enumerate(prompts, 1): | |
| print("=" * 80) | |
| print(f"PROMPT {i}: {prompt}") | |
| print("-" * 80) | |
| print(generate_text(model, tokenizer, prompt, device=device)) | |
| print("=" * 80) | |
| if __name__ == "__main__": | |
| main() | |