Files
opencv/samples/dnn/bert_inference.py
T
Jaivardhan Bhola 84c2360c27 Merge pull request #29675 from jaivardhan-bhola:generalized-tokenizer
Generalize tokenizer loading to support method-based family dispatch in DNN - #29675

Companion PR: https://github.com/opencv/opencv_extra/pull/1402

### Changes
Added ALBERT and BERT support end-to-end, with `samples/dnn/albert_inference.py `and `samples/dnn/bert_inference.py `as validation samples, plus expanded coverage in modules/dnn/test/test_tokenizer.cpp. Required changes to `cv::dnn::dnn.hpp`, `graph_fusion_attention.cpp`, and `unicode.cpp/unicode.hpp` to support Unigram and WordPiece tokenizers.

To back ALBERT/BERT, generalized `cv::dnn::Tokenizer `from a single BPE implementation into a method-dispatched frontend, adding `core_wordpiece.cpp/hpp` (WordPiece) and `core_unigram.cpp/hpp` (Unigram) as new backends. `tokenizer.cpp` now routes by method across BPE, Gemma,, SentencePiece, Unigram, and WordPiece behind one shared interface.

Tested against the following samples and the output matches to old tokenizer:

```
gpt2_inference.py
qwen_inference.py
gemma3_inference.py
```
GPT2:

```
Preparing GPT-2 model...
Inferencing GPT-2 model...
Hello, I'm a language model, not a programming language. I'm a language model. I'm a language model. I'm a language model. I'm a language model. I'm a
```

Gemma3:
```
Preparing Gemma3 model...
Prompt:
<start_of_turn>user
What is OpenCV?<end_of_turn>
<start_of_turn>model

Inferencing Gemma3 model...
Response:
Okay, let's break down what OpenCV is.

**What is OpenCV?**

OpenCV (Open Source Computer Vision Library) is a powerful and
```

Qwen2.5:
```
Preparing Qwen2.5 model...
Prompt:
<|im_start|>user
What is OpenCV?<|im_end|>
<|im_start|>assistant

Inferencing Qwen2.5 model...
Response:
OpenCV is a set of computer vision libraries in C++ designed to be used for image and video processing. It provides a wide range of tools and functions for
```

### Tokenizer References 
Byte-level BPE: [tokenizers/src/pre_tokenizers/byte_level.rs](https://github.com/huggingface/tokenizers/blob/main/tokenizers/src/pre_tokenizers/byte_level.rs)
(This defines the byte-level mapping rules, which is used in conjunction with the [BPE model](https://www.google.com/search?q=https://github.com/huggingface/tokenizers/blob/main/tokenizers/src/models/bpe/mod.rs))

SentencePiece BPE (Metaspace): [tokenizers/src/pre_tokenizers/metaspace.rs](https://www.google.com/search?q=https://github.com/huggingface/tokenizers/blob/main/tokenizers/src/pre_tokenizers/metaspace.rs)
(This defines the rule for replacing whitespace with the U+2581 _ character and handling byte fallback)

Unigram: [tokenizers/src/models/unigram/mod.rs](https://www.google.com/search?q=https://github.com/huggingface/tokenizers/blob/main/tokenizers/src/models/unigram/mod.rs)
(This contains the core logic for the Unigram lattice scoring and probabilistic tokenization rules)

WordPiece: [tokenizers/src/models/wordpiece/mod.rs](https://www.google.com/search?q=https://github.com/huggingface/tokenizers/blob/main/tokenizers/src/models/wordpiece/mod.rs)
(This explicitly cites Schuster & Nakajima in the code comments and implements the greedy longest-match rule with the ## prefix)

### Pull Request Readiness Checklist

- [x] I agree to contribute to the project under Apache 2 License.
- [x] To the best of my knowledge, the proposed patch is not based on
      code under GPL or another license incompatible with OpenCV.
- [x] The PR is proposed to the proper branch (`5.x`).
- [x] There is a reference to the original bug report and related work.
- [x] There is accuracy test and test data in `opencv_extra`, same
      branch name (`generalized-tokenizer`) — `bert/`, `t5/` fixtures
      back the new C++ tests.
- [x] The feature is documented and sample code builds with project CMake.
2026-09-23 16:20:33 +03:00

102 lines
3.6 KiB
Python

# This file is part of OpenCV project.
# It is subject to the license terms in the LICENSE file found in the top-level directory
# of this distribution and at http://opencv.org/license.html.
# Copyright (C) 2026, BigVision LLC, all rights reserved.
# Third party copyrights are property of their respective owners.
'''
This is a sample script to run BERT (bert-base-uncased) masked-LM inference
in OpenCV using an ONNX model. The input text must contain a single literal
"[MASK]" token; the script prints the top predictions for that position.
Model: https://huggingface.co/google-bert/bert-base-uncased
Downloading the BERT model and tokenizer:
1. Install the Hugging Face CLI:
pip install -U "hf"
2. Download only the files needed (model.onnx, config.json and the WordPiece
tokenizer.json/vocab.txt) into a local directory:
hf download google-bert/bert-base-uncased \
model.onnx config.json tokenizer.json tokenizer_config.json vocab.txt \
--local-dir bert-base-uncased
Run the script:
1. Install the required dependencies:
pip install numpy
2. Run the script:
python bert_inference.py --model=<path-to-onnx-model> \
--tokenizer_path=<path-to-bert-base-uncased-dir> \
--text="Paris is the [MASK] of France." \
--topk=5
'''
import numpy as np
import argparse
import cv2 as cv
MASK_TOKEN = '[MASK]'
def parse_args():
parser = argparse.ArgumentParser(description='Use this script to run BERT masked-LM inference in OpenCV',
formatter_class=argparse.ArgumentDefaultsHelpFormatter)
parser.add_argument('--model', type=str, required=True, help='Path to the BERT ONNX model file.')
parser.add_argument('--tokenizer_path', type=str, required=True, help='Path to the BERT tokenizer directory, or to its config.json.')
parser.add_argument('--text', type=str, default='Paris is the [MASK] of France.', help='Input text containing a single [MASK] token.')
parser.add_argument('--topk', type=int, default=5, help='Number of top predictions to print.')
return parser.parse_args()
def encode_with_mask(tokenizer, text):
if text.count(MASK_TOKEN) != 1:
raise ValueError('expected exactly one [MASK] token')
mask_id = tokenizer.encode(MASK_TOKEN)[1]
ids = list(tokenizer.encode(text))
mask_pos = ids.index(mask_id)
return ids, mask_pos
def softmax(x):
e = np.exp(x - np.max(x))
return e / e.sum()
def bert_inference(net, tokenizer, text, topk):
print("Inferencing BERT model...")
ids, mask_pos = encode_with_mask(tokenizer, text)
n = len(ids)
input_ids = np.array([ids], dtype=np.int64)
attention_mask = np.ones((1, n), dtype=np.int64)
token_type_ids = np.zeros((1, n), dtype=np.int64)
net.setInput(input_ids, 'input_ids')
net.setInput(attention_mask, 'attention_mask')
net.setInput(token_type_ids, 'token_type_ids')
logits = net.forward('logits')
row = logits[0, mask_pos]
probs = softmax(row)
top_ids = np.argsort(row)[::-1][:topk]
return [(int(t), tokenizer.decode([int(t)]).strip(), float(probs[t])) for t in top_ids]
if __name__ == '__main__':
args = parse_args()
print("Preparing BERT model...")
tokenizer = cv.dnn.Tokenizer.load(args.tokenizer_path)
net = cv.dnn.readNetFromONNX(args.model, cv.dnn.ENGINE_OPENCV)
print(f"Text: {args.text}")
predictions = bert_inference(net, tokenizer, args.text, args.topk)
for rank, (token_id, token, prob) in enumerate(predictions, start=1):
print(f" {rank}. {token!r} id={token_id} p={prob:.4f}")