Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -289,9 +289,15 @@ Useful options:
```shell
--model_dir baidu/Unlimited-OCR # Local path or Hugging Face model ID
--gpu 0 # CUDA_VISIBLE_DEVICES value
--repetition_penalty 1.1 # Repetition penalty factor to prevent text looping (>1.0 to enable)
--no_repeat_ngram_size 35 # N-gram size to ban repeated sequences (0 to disable)
--ngram_window 128 # Sliding window size for n-gram repetition blocking
--server_log ./log/sglang_server.log
```

> **Preventing Text Repetition Loops:**
> On certain dense or structured document pages, short text phrases may repeat endlessly. Pass `--repetition_penalty 1.1` in `infer.py` or set `repetition_penalty=1.1` in `model.generate()` to prevent repetition loops. You can also tune `no_repeat_ngram_size` (e.g. `15`-`35`) to detect shorter repeating units.

For OmniDocBench evaluation, you need to perform the following post-processing.
```python
def remove_det(raw: str) -> str:
Expand Down
14 changes: 11 additions & 3 deletions infer.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
CONTEXT_LENGTH = 32768
NO_REPEAT_NGRAM_SIZE = 35
NGRAM_WINDOW = 128
REPETITION_PENALTY = 1.0
REQUEST_TIMEOUT = 1200
MAX_RETRIES = 5
NO_REPEAT_NGRAM_PROCESSOR_STR = None
Expand Down Expand Up @@ -194,11 +195,15 @@ def infer_one(image_path: str, output_file: str | None, args, idx: int) -> dict:
"stream": True,
"images_config": {"image_mode": args.image_mode},
}
if NO_REPEAT_NGRAM_SIZE > 0 and NGRAM_WINDOW > 0:
if getattr(args, "repetition_penalty", REPETITION_PENALTY) > 1.0:
payload["repetition_penalty"] = args.repetition_penalty
ngram_size = getattr(args, "no_repeat_ngram_size", NO_REPEAT_NGRAM_SIZE)
window_size = getattr(args, "ngram_window", NGRAM_WINDOW)
if ngram_size > 0 and window_size > 0:
payload["custom_logit_processor"] = get_ngram_processor_str()
payload["custom_params"] = {
"ngram_size": NO_REPEAT_NGRAM_SIZE,
"window_size": NGRAM_WINDOW,
"ngram_size": ngram_size,
"window_size": window_size,
}

name = os.path.basename(image_path)
Expand Down Expand Up @@ -312,6 +317,9 @@ def parse_args():
parser.add_argument("--gpu", default="0")
parser.add_argument("--model_dir", default="baidu/Unlimited-OCR")
parser.add_argument("--image_mode", choices=("gundam", "base"), default="gundam")
parser.add_argument("--no_repeat_ngram_size", type=int, default=35, help="N-gram size to ban repeated sequences (0 to disable)")
parser.add_argument("--ngram_window", type=int, default=128, help="Sliding window size for n-gram repetition blocking")
parser.add_argument("--repetition_penalty", type=float, default=1.0, help="Repetition penalty factor to prevent text looping (>1.0 to enable)")
parser.add_argument("--server_log", default="./log/sglang_server.log")
return parser.parse_args()

Expand Down
66 changes: 66 additions & 0 deletions tests/test_infer_repetition.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
"""
Unit test for repetition penalty and ngram parameters in infer.py (Fixes #55).
"""

import sys
import unittest
from unittest.mock import MagicMock, patch

import infer


class TestInferRepetitionControl(unittest.TestCase):
def test_parse_args_defaults(self):
with patch.object(sys, "argv", ["infer.py"]):
args = infer.parse_args()
self.assertEqual(args.repetition_penalty, 1.0)
self.assertEqual(args.no_repeat_ngram_size, 35)
self.assertEqual(args.ngram_window, 128)

def test_parse_args_custom_repetition(self):
test_argv = [
"infer.py",
"--repetition_penalty",
"1.1",
"--no_repeat_ngram_size",
"15",
"--ngram_window",
"256",
]
with patch.object(sys, "argv", test_argv):
args = infer.parse_args()
self.assertEqual(args.repetition_penalty, 1.1)
self.assertEqual(args.no_repeat_ngram_size, 15)
self.assertEqual(args.ngram_window, 256)

@patch("infer.requests.post")
@patch("infer.get_ngram_processor_str", return_value="MockNGramProcessor")
@patch("infer.build_content", return_value=[{"type": "text", "text": "test"}])
def test_infer_one_payload_building(self, mock_build_content, mock_get_processor, mock_post):
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.iter_lines.return_value = [b"data: [DONE]"]
mock_post.return_value = mock_response

args = MagicMock()
args.image_mode = "gundam"
args.repetition_penalty = 1.1
args.no_repeat_ngram_size = 20
args.ngram_window = 64

infer.infer_one("dummy.png", None, args, 1)

self.assertTrue(mock_post.called)
_, kwargs = mock_post.call_args
import json

payload = json.loads(kwargs["data"])

self.assertEqual(payload.get("repetition_penalty"), 1.1)
self.assertEqual(payload.get("custom_logit_processor"), "MockNGramProcessor")
self.assertEqual(payload["custom_params"]["ngram_size"], 20)
self.assertEqual(payload["custom_params"]["window_size"], 64)


if __name__ == "__main__":
unittest.main()