Views
No views yet
| model | token_acc ↑ | normed edit distance ↓ |
|---|---|---|
| pix2tex | 0.5346 | 0.10312 |
| pix2tex* | 0.60 | 0.10 |
| nougat-latex-based | 0.623850 | 0.06180 |
pip install transformers >= 4.34.0The inference API widget sometimes cuts the response short. Please check this issue for more details. You may want to run the model yourself in case the inference API bug cuts the results short.
1git clone git@github.com:NormXU/nougat-latex-ocr.git
2cd ./nougat-latex-ocr1import torch
2from PIL import Image
3from transformers import VisionEncoderDecoderModel
4from transformers.models.nougat import NougatTokenizerFast
5from nougat_latex import NougatLaTexProcessor
6
7model_name = "Norm/nougat-latex-base"
8device = "cuda" if torch.cuda.is_available() else "cpu"
9# init model
10model = VisionEncoderDecoderModel.from_pretrained(model_name).to(device)
11
12# init processor
13tokenizer = NougatTokenizerFast.from_pretrained(model_name)
14
15latex_processor = NougatLaTexProcessor.from_pretrained(model_name)
16
17# run test
18image = Image.open("path/to/latex/image.png")
19if not image.mode == "RGB":
20 image = image.convert('RGB')
21
22pixel_values = latex_processor(image, return_tensors="pt").pixel_values
23
24decoder_input_ids = tokenizer(tokenizer.bos_token, add_special_tokens=False,
25 return_tensors="pt").input_ids
26with torch.no_grad():
27 outputs = model.generate(
28 pixel_values.to(device),
29 decoder_input_ids=decoder_input_ids.to(device),
30 max_length=model.decoder.config.max_length,
31 early_stopping=True,
32 pad_token_id=tokenizer.pad_token_id,
33 eos_token_id=tokenizer.eos_token_id,
34 use_cache=True,
35 num_beams=5,
36 bad_words_ids=[[tokenizer.unk_token_id]],
37 return_dict_in_generate=True,
38 )
39sequence = tokenizer.batch_decode(outputs.sequences)[0]
40sequence = sequence.replace(tokenizer.eos_token, "").replace(tokenizer.pad_token, "").replace(tokenizer.bos_token, "")
41print(sequence)
42