Spaces:
Sleeping
Sleeping
| 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) | |