jackbean/question_generation_api
0
1 2from fastapi import FastAPI, HTTPException3from pydantic import BaseModel4import torch5from transformers import T5ForConditionalGeneration, T5Tokenizer6 7# Initialize FastAPI app8app = FastAPI()9 10# Lazy load model and tokenizer11model = None12tokenizer = None13 14def load_model():15 global model, tokenizer16 if model is None or tokenizer is None:17 tokenizer = T5Tokenizer.from_pretrained('./tokenizer12')18 model = T5ForConditionalGeneration.from_pretrained('./model')19 model.to('cuda' if torch.cuda.is_available() else 'cpu')20 21# Request body schema using Pydantic22class QuestionRequest(BaseModel):23 context: str24 answer: str25 26from fastapi import Query27 28@app.post("/generate_question")29async def generate_question(request: QuestionRequest):30 load_model()31 device = 'cuda' if torch.cuda.is_available() else 'cpu'32 33 input_text = f"context: {request.context} answer: {request.answer}"34 encoding = tokenizer.encode_plus(35 input_text,36 max_length=512,37 padding="max_length",38 truncation=True,39 return_tensors="pt"40 )41 input_ids = encoding["input_ids"].to(device)42 attention_mask = encoding["attention_mask"].to(device)43 44 model.eval()45 with torch.no_grad():46 beam_outputs = model.generate(47 input_ids=input_ids,48 attention_mask=attention_mask,49 max_length=72,50 early_stopping=True,51 num_beams=5,52 num_return_sequences=353 )54 55 return {56 "generated_questions": [57 tokenizer.decode(output, skip_special_tokens=True, clean_up_tokenization_spaces=True)58 for output in beam_outputs59 ]60 }61 62 63 64if __name__ == "__main__":65 import uvicorn66 uvicorn.run(app, host="0.0.0.0", port=7860)67 68 