import io import litserve as ls from transformers import NougatProcessor, VisionEncoderDecoderModel import torch from PIL import Image from fastapi import UploadFile from pdf2image import convert_from_bytes class NougatLitAPI(ls.LitAPI): def setup(self, device): model_name = "facebook/nougat-base" self.processor = NougatProcessor.from_pretrained(model_name) self.model = VisionEncoderDecoderModel.from_pretrained(model_name) self.model.to(device) self.model.eval() def decode_request(self, request: UploadFile): file_bytes = request.file.read() # convert PDF bytes to images try: images = convert_from_bytes(file_bytes, dpi=300) images = [img.convert("RGB") for img in images] except Exception as e: print(f"couldn't convert pdf to images: {e}") images = [] if not images: # fallback: blank white image if PDF fails img = Image.new('RGB', (500, 500), color="white") images.append(img) print("warning: no images found in pdf, using blank fallback.") return images def predict(self, images): markdown_output = [] for i, img in enumerate(images): # preprocess single page pixel_values = self.processor(images=img, return_tensors="pt").pixel_values if torch.cuda.is_available(): pixel_values = pixel_values.to(self.model.device) # generate markdown for page with torch.no_grad(): outputs = self.model.generate( pixel_values, min_length=1, max_new_tokens=2048, # enough for one page bad_words_ids=[[self.processor.tokenizer.unk_token_id]], ) decoded_text = self.processor.batch_decode(outputs, skip_special_tokens=True)[0] decoded_text = self.processor.post_process_generation(decoded_text, fix_markdown=False) markdown_output.append(f"\n\n--- Page {i + 1} ---\n\n{decoded_text}") return "".join(markdown_output) def encode_response(self, output_text): return {"markdown_output": output_text} if __name__ == "__main__": api = NougatLitAPI() server = ls.LitServer(api, accelerator="cuda", devices=1) server.run(port=8000)