mirror of
https://github.com/opencv/opencv.git
synced 2026-10-06 04:03:36 +03:00
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.
102 lines
3.6 KiB
Python
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}")
|