From 37728eed3409640bc06a06aa1b1c8ecca68d80ce Mon Sep 17 00:00:00 2001 From: not-knope <102312680+not-knope@users.noreply.github.com> Date: Tue, 21 Jul 2026 22:11:25 +0200 Subject: [PATCH] fix(infer): add repetition_penalty and CLI ngram parameters to prevent text looping (Fixes #55) --- README.md | 6 ++++ infer.py | 14 ++++++-- tests/test_infer_repetition.py | 66 ++++++++++++++++++++++++++++++++++ 3 files changed, 83 insertions(+), 3 deletions(-) create mode 100644 tests/test_infer_repetition.py diff --git a/README.md b/README.md index a3ceffe..62979b5 100644 --- a/README.md +++ b/README.md @@ -286,9 +286,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. + ## Visualization diff --git a/infer.py b/infer.py index a52ab9e..bb9418c 100644 --- a/infer.py +++ b/infer.py @@ -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 @@ -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) @@ -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() diff --git a/tests/test_infer_repetition.py b/tests/test_infer_repetition.py new file mode 100644 index 0000000..3528443 --- /dev/null +++ b/tests/test_infer_repetition.py @@ -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()