Skip to content

perf: true batched inference for TransformerDetector (fixes #23) - #113

Open
c71qu3 wants to merge 3 commits into
KRLabsOrg:mainfrom
c71qu3:feat/transformer-batch-processing
Open

c71qu3 wants to merge 3 commits into
KRLabsOrg:mainfrom
c71qu3:feat/transformer-batch-processing

Conversation

@c71qu3

@c71qu3 c71qu3 commented Sep 17, 2026 •

Copy link
Copy Markdown

TransformerDetector.predict_prompt_batch now performs padded batch tokenization and a single model forward pass per configured batch_size, preserving input order and correctly trimming prompt/padding per sample for both tokens and spans. Add strict len(prompts) == len(answers) validation (no silent zip truncation), including in the LLM detector request path.

Tests cover uneven sequence lengths, batch_size 1 and >1, output order, tokens, spans, min_confidence filtering, empty input, and mismatch errors; a spy/stub verifies one forward call per transformer batch without downloading a model.

Summary

  • Implements true batched inference in TransformerDetector.predict_prompt_batch: padded batch tokenization + one model forward pass per batch (configurable batch_size), preserving input order and trimming prompt/padding per sample correctly for both tokens and spans.
  • Adds strict input validation: raises ValueError when len(prompts) != len(answers) (no silent truncation) and rejects invalid batch_size < 1; keeps the LLM detector’s concurrent-request path but applies the same length validation.
  • Ensures behavior parity with the single-example path: output matches predict_prompt per sample (within tolerance), including min_confidence filtering and taxonomy typing when configured.
  • Adds/updates tests to cover uneven sequence lengths, batch size 1 and >1, stable output order, empty input, mismatch errors, and a spy/stub model to verify one forward call per transformer batch without downloading a model.

Related issue

Closes #23.

Type of change

  • Bug fix
  • Feature
  • Documentation
  • Tests
  • Refactor or maintenance

Testing

  • ruff format --check lettucedetect/ lettucedetect_api/ tests/
  • ruff check lettucedetect/ lettucedetect_api/ tests/ --extend-exclude lettucedetect/integrations/
  • python -m pytest
  • Other:

Checklist

  • I kept the PR focused on one change.
  • I added or updated tests/docs when needed.
  • I checked that no secrets, API keys, or credentials are included.

Rights & sign-off (required)

  • I certify that I have the right to submit this code and that it may be
    distributed under the repository's MIT license
    (see CONTRIBUTING).

)

TransformerDetector.predict_prompt_batch now performs padded batch tokenization
and a single model forward pass per configured batch_size, preserving input
order and correctly trimming prompt/padding per sample for both tokens and spans.
Add strict len(prompts) == len(answers) validation (no silent zip truncation),
including in the LLM detector request path.

Tests cover uneven sequence lengths, batch_size 1 and >1, output order, tokens,
spans, min_confidence filtering, empty input, and mismatch errors; a spy/stub
verifies one forward call per transformer batch without downloading a model.

@adaamko adaamko left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for this, @c71qu3. This is the first implementation for #23 and it's the one we'd like to merge. I checked it against the real model: KRLabsOrg/lettucedect-base-modernbert-en-v1 on CPU, 7 prompt/answer pairs from 67 to ~3,200 tokens, batch sizes 1/2/3/7. Batched and one-by-one spans match exactly, and token probabilities agree within 2e-6 (exactly with attn_implementation="eager"). The mmBERT + taxonomy-head cascade also gives identical typed spans, and the truncation path matches too. Using sequence_ids for the per-row answer start and attention_mask for the real length is the right call. A few things before merge:

  1. Lint/format (merge gate). ruff format --check lettucedetect/ lettucedetect_api/ tests/ wants to reformat lettucedetect/detectors/transformer.py and tests/test_inference_pytest.py. ruff check lettucedetect/ lettucedetect_api/ tests/ --extend-exclude lettucedetect/integrations/ reports 12 errors:

    • D205/D209 docstring layout at transformer.py:273 and tests/test_inference_pytest.py:804, 833, 876, 910
    • D107/D102 on CountingClient in tests/test_llm_detector_pytest.py:35,39
    • S105 at tests/test_inference_pytest.py:819, a false positive on the "[PAD]" token string; a # noqa: S105 is fine

    Running ruff format and then ruff check --fix handles most of it.

  2. Bounded default batch size. Before this change predict_prompt_batch used constant memory whatever the list length. With batch_size=None meaning "whole list" (transformer.py:560), existing callers that pass a few thousand pairs now get one huge forward pass and can run out of memory. Please default to a fixed size (for example batch_size: int = 16 in __init__) and keep the per-call override.

  3. Validate batch_size. A negative value currently returns [] silently (range(0, n, -2) is empty), and 0 falls back to the default via or. Raise ValueError for anything below 1.

  4. PR description. Please fill in the Summary section, add Closes #23 under "Related issue", and tick the ruff boxes once they pass.

Non-blocking nits:

  • tests/test_inference_pytest.py:526-527: delete the commented-out tokenizer.return_value line and the stale comment above it.
  • transformer.py:282: instead of rejecting left-padding tokenizers, you could pass padding_side="right" to the tokenizer call in _predict_batch, which is supported at our transformers>=4.48.3 floor. If you keep the check, the message reads better as "...requires a right-padding tokenizer".
  • transformer.py:319-323: this warning fires whenever seq_len == max_length, even for inputs that fit exactly, and it uses an f-string where the rest of the module uses %-style logger args. Measuring the untruncated length, like predict_prompt does, would match the single-path warning.
  • Typos in test docstrings: "modal" → "model", "everthing" → "everything".
  • Optional follow-up: sorting pairs by length before micro-batching (restoring order afterwards) would cut padding waste on mixed-length inputs.

Once 1–4 are in and CI is green, we'll merge. Thanks again!

@c71qu3
c71qu3 force-pushed the feat/transformer-batch-processing branch from a1dd2b7 to 3d2a380 Compare September 29, 2026 12:04
@c71qu3

c71qu3 commented Sep 29, 2026

Copy link
Copy Markdown
Author

Thanks for the detailed review!

I’ve pushed updates to address the requested changes.

@c71qu3
c71qu3 requested a review from adaamko September 29, 2026 12:09

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Support for “real” batch processing

2 participants