DeepDigest / server.py
ArielKes's picture
app and server code. app sends pdf to server that returns mmd
64cf084
Raw
History Blame Contribute Delete
2.4 kB
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)