Add OCL support to improve performance

This commit is contained in:
WarmUpTill
2023-05-24 14:56:35 +02:00
committed by WarmUpTill
parent f506600b8c
commit 284c8020b2
6 changed files with 83 additions and 52 deletions

View File

@@ -83,6 +83,11 @@ const static std::map<tesseract::PageSegMode, std::string> pageSegModes = {
"AdvSceneSwitcher.condition.video.ocrMode.sparseTextOSD"},
};
MacroConditionVideo::MacroConditionVideo(Macro *m) : MacroCondition(m, true)
{
SetupOpenCL();
}
cv::CascadeClassifier initObjectCascade(std::string &path)
{
cv::CascadeClassifier cascade;
@@ -261,18 +266,21 @@ bool MacroConditionVideo::SetLanguage(const std::string &language)
bool MacroConditionVideo::ScreenshotContainsPattern()
{
cv::Mat result;
cv::UMat result;
MatchPattern(_screenshotData.image, _patternImageData,
_patternMatchParameters.threshold, result,
_patternMatchParameters.useAlphaAsMask,
_patternMatchParameters.matchMode);
if (result.total() == 0) {
return false;
}
return countNonZero(result) > 0;
}
bool MacroConditionVideo::OutputChanged()
{
if (_patternMatchParameters.useForChangedCheck) {
cv::Mat result;
cv::UMat result;
_patternImageData = CreatePatternData(_matchImage);
MatchPattern(_screenshotData.image, _patternImageData,
_patternMatchParameters.threshold, result,

View File

@@ -25,7 +25,7 @@ class PreviewDialog;
class MacroConditionVideo : public MacroCondition {
public:
MacroConditionVideo(Macro *m) : MacroCondition(m, true){};
MacroConditionVideo(Macro *m);
bool CheckCondition();
bool Save(obs_data_t *obj) const;
bool Load(obs_data_t *obj);

View File

@@ -1,8 +1,12 @@
#include "opencv-helpers.hpp"
#include "log-helper.hpp"
#include <opencv2/core/ocl.hpp>
#include <opencv2/core/mat.hpp>
namespace advss {
PatternImageData CreatePatternData(QImage &pattern)
PatternImageData CreatePatternData(const QImage &pattern)
{
PatternImageData data{};
if (pattern.isNull()) {
@@ -20,18 +24,20 @@ PatternImageData CreatePatternData(QImage &pattern)
return data;
}
static void invertPatternMatchResult(cv::Mat &mat)
static void invertPatternMatchResult(cv::UMat &mat)
{
for (int r = 0; r < mat.rows; r++) {
for (int c = 0; c < mat.cols; c++) {
float value = mat.at<float>(r, c) =
1.0 - mat.at<float>(r, c);
auto temp = mat.getMat(cv::ACCESS_RW);
for (int r = 0; r < temp.rows; r++) {
for (int c = 0; c < temp.cols; c++) {
float value = temp.at<float>(r, c) =
1.0 - temp.at<float>(r, c);
}
}
mat = temp.getUMat(cv::ACCESS_RW);
}
void MatchPattern(QImage &img, const PatternImageData &patternData,
double threshold, cv::Mat &result, bool useAlphaAsMask,
double threshold, cv::UMat &result, bool useAlphaAsMask,
cv::TemplateMatchModes matchMode)
{
if (img.isNull() || patternData.rgbaPattern.empty()) {
@@ -50,13 +56,12 @@ void MatchPattern(QImage &img, const PatternImageData &patternData,
// thus should not be used while matching the pattern as well
//
// Input format is Format_RGBA8888 so discard the 4th channel
std::vector<cv::Mat1b> inputChannels;
std::vector<cv::UMat> inputChannels;
cv::split(input, inputChannels);
std::vector<cv::Mat1b> rgbChanlesImage(
std::vector<cv::UMat> rgbChanlesImage(
inputChannels.begin(), inputChannels.begin() + 3);
cv::Mat3b rgbInput;
cv::UMat rgbInput;
cv::merge(rgbChanlesImage, rgbInput);
cv::matchTemplate(rgbInput, patternData.rgbPattern, result,
matchMode, patternData.mask);
} else {
@@ -75,7 +80,7 @@ void MatchPattern(QImage &img, const PatternImageData &patternData,
}
void MatchPattern(QImage &img, QImage &pattern, double threshold,
cv::Mat &result, bool useAlphaAsMask,
cv::UMat &result, bool useAlphaAsMask,
cv::TemplateMatchModes matchColor)
{
auto data = CreatePatternData(pattern);
@@ -91,9 +96,9 @@ std::vector<cv::Rect> MatchObject(QImage &img, cv::CascadeClassifier &cascade,
return {};
}
auto i = QImageToMat(img);
cv::Mat frameGray;
cv::cvtColor(i, frameGray, cv::COLOR_RGBA2GRAY);
auto image = QImageToMat(img);
cv::UMat frameGray;
cv::cvtColor(image, frameGray, cv::COLOR_RGBA2GRAY);
cv::equalizeHist(frameGray, frameGray);
std::vector<cv::Rect> objects;
cascade.detectMultiScale(frameGray, objects, scaleFactor, minNeighbors,
@@ -120,7 +125,7 @@ uchar GetAvgBrightness(QImage &img)
return brightnessSum / (hsvImage.rows * hsvImage.cols);
}
cv::Mat PreprocessForOCR(const QImage &image, const QColor &color)
cv::UMat PreprocessForOCR(const QImage &image, const QColor &color)
{
auto mat = QImageToMat(image);
@@ -160,7 +165,8 @@ std::string RunOCR(tesseract::TessBaseAPI *ocr, const QImage &image,
#ifdef OCR_SUPPORT
auto mat = PreprocessForOCR(image, color);
ocr->SetImage(mat.data, mat.cols, mat.rows, 1, mat.step);
ocr->SetImage(mat.getMat(cv::ACCESS_READ).data, mat.cols, mat.rows, 1,
mat.step);
ocr->Recognize(0);
std::unique_ptr<char[]> detectedText(ocr->GetUTF8Text());
@@ -207,13 +213,14 @@ bool ContainsPixelsInColorRange(const QImage &image, const QColor &color,
// Assumption is that QImage uses Format_RGBA8888.
// Conversion from: https://github.com/dbzhang800/QtOpenCV
cv::Mat QImageToMat(const QImage &img)
cv::UMat QImageToMat(const QImage &img)
{
if (img.isNull()) {
return cv::Mat();
return cv::UMat();
}
return cv::Mat(img.height(), img.width(), CV_8UC(img.depth() / 8),
(uchar *)img.bits(), img.bytesPerLine());
auto temp = cv::Mat(img.height(), img.width(), CV_8UC(img.depth() / 8),
(uchar *)img.bits(), img.bytesPerLine());
return temp.getUMat(cv::ACCESS_RW);
}
QImage MatToQImage(const cv::Mat &mat)
@@ -225,4 +232,12 @@ QImage MatToQImage(const cv::Mat &mat)
QImage::Format::Format_RGBA8888);
}
void SetupOpenCL()
{
if (cv::ocl::haveOpenCL() && !cv::ocl::useOpenCL()) {
blog(LOG_INFO, "enabled OpenCL support for OpenCV");
cv::ocl::setUseOpenCL(true);
}
}
} // namespace advss

View File

@@ -42,29 +42,30 @@ constexpr int maxMinNeighbors = 6;
constexpr double defaultScaleFactor = 1.1;
struct PatternImageData {
cv::Mat4b rgbaPattern;
cv::Mat3b rgbPattern;
cv::Mat1b mask;
cv::UMat rgbaPattern;
cv::UMat rgbPattern;
cv::UMat mask;
};
PatternImageData CreatePatternData(QImage &pattern);
PatternImageData CreatePatternData(const QImage &pattern);
void MatchPattern(QImage &img, const PatternImageData &patternData,
double threshold, cv::Mat &result, bool useAlphaAsMask,
double threshold, cv::UMat &result, bool useAlphaAsMask,
cv::TemplateMatchModes matchMode);
void MatchPattern(QImage &img, QImage &pattern, double threshold,
cv::Mat &result, bool useAlphaAsMask,
cv::UMat &result, bool useAlphaAsMask,
cv::TemplateMatchModes matchMode);
std::vector<cv::Rect> MatchObject(QImage &img, cv::CascadeClassifier &cascade,
double scaleFactor, int minNeighbors,
const cv::Size &minSize,
const cv::Size &maxSize);
uchar GetAvgBrightness(QImage &img);
cv::Mat PreprocessForOCR(const QImage &image, const QColor &color);
cv::UMat PreprocessForOCR(const QImage &image, const QColor &color);
std::string RunOCR(tesseract::TessBaseAPI *, const QImage &, const QColor &);
bool ContainsPixelsInColorRange(const QImage &image, const QColor &color,
double colorDeviationThreshold,
double totalPixelMatchThreshold);
cv::Mat QImageToMat(const QImage &img);
cv::UMat QImageToMat(const QImage &img);
QImage MatToQImage(const cv::Mat &mat);
void SetupOpenCL();
} // namespace advss

View File

@@ -122,7 +122,6 @@ void PreviewDialog::PatternMatchParametersChanged(
{
std::unique_lock<std::mutex> lock(_mtx);
_patternMatchParams = params;
_patternImageData = CreatePatternData(_patternMatchParams.image);
}
void PreviewDialog::ObjDetectParametersChanged(const ObjDetectParameters &params)
@@ -170,8 +169,8 @@ void PreviewDialog::UpdateImage(const QPixmap &image)
if (_type == PreviewType::SELECT_AREA && !_selectingArea) {
DrawFrame();
}
emit NeedImage(_video, _type, _patternMatchParams, _patternImageData,
_objDetectParams, _ocrParams, _areaParams, _condition);
emit NeedImage(_video, _type, _patternMatchParams, _objDetectParams,
_ocrParams, _areaParams, _condition);
}
void PreviewDialog::Start()
@@ -187,7 +186,7 @@ void PreviewDialog::Start()
return;
}
PreviewImage *worker = new PreviewImage();
PreviewImage *worker = new PreviewImage(_mtx);
worker->moveToThread(&_thread);
connect(&_thread, &QThread::finished, worker, &QObject::deleteLater);
connect(worker, &PreviewImage::ImageReady, this,
@@ -198,8 +197,8 @@ void PreviewDialog::Start()
&PreviewImage::CreateImage);
_thread.start();
emit NeedImage(_video, _type, _patternMatchParams, _patternImageData,
_objDetectParams, _ocrParams, _areaParams, _condition);
emit NeedImage(_video, _type, _patternMatchParams, _objDetectParams,
_ocrParams, _areaParams, _condition);
}
void PreviewDialog::DrawFrame()
@@ -217,13 +216,14 @@ void PreviewDialog::DrawFrame()
_rubberBand->show();
}
static void markPatterns(cv::Mat &matchResult, QImage &image,
const cv::Mat &pattern)
static void markPatterns(cv::UMat &matchResult, QImage &image,
const cv::UMat &pattern)
{
auto temp = matchResult.getMat(cv::ACCESS_RW);
auto matchImg = QImageToMat(image);
for (int row = 0; row < matchResult.rows - 1; row++) {
for (int col = 0; col < matchResult.cols - 1; col++) {
if (matchResult.at<float>(row, col) != 0.0) {
for (int row = 0; row < temp.rows - 1; row++) {
for (int col = 0; col < temp.cols - 1; col++) {
if (temp.at<float>(row, col) != 0.0) {
rectangle(matchImg, {col, row},
cv::Point(col + pattern.cols,
row + pattern.rows),
@@ -231,6 +231,7 @@ static void markPatterns(cv::Mat &matchResult, QImage &image,
}
}
}
matchResult = temp.getUMat(cv::ACCESS_RW);
}
static void markObjects(QImage &image, std::vector<cv::Rect> &objects)
@@ -244,9 +245,10 @@ static void markObjects(QImage &image, std::vector<cv::Rect> &objects)
}
}
PreviewImage::PreviewImage(std::mutex &mtx) : _mtx(mtx) {}
void PreviewImage::CreateImage(const VideoInput &video, PreviewType type,
const PatternMatchParameters &patternMatchParams,
const PatternImageData &patternImageData,
ObjDetectParameters objDetectParams,
OCRParameters ocrParams,
const AreaParameters &areaParams,
@@ -271,11 +273,14 @@ void PreviewImage::CreateImage(const VideoInput &video, PreviewType type,
}
if (type == PreviewType::SHOW_MATCH) {
std::unique_lock<std::mutex> lock(_mtx);
if (areaParams.enable) {
screenshot.image = screenshot.image.copy(
areaParams.area.x, areaParams.area.y,
areaParams.area.width, areaParams.area.height);
}
const auto patternImageData =
CreatePatternData(patternMatchParams.image);
// Will emit status label update
MarkMatch(screenshot.image, patternMatchParams,
patternImageData, objDetectParams, ocrParams,
@@ -294,12 +299,12 @@ void PreviewImage::MarkMatch(QImage &screenshot,
VideoCondition condition)
{
if (condition == VideoCondition::PATTERN) {
cv::Mat result;
cv::UMat result;
MatchPattern(screenshot, patternImageData,
patternMatchParams.threshold, result,
patternMatchParams.useAlphaAsMask,
patternMatchParams.matchMode);
if (countNonZero(result) == 0) {
if (result.total() == 0 || countNonZero(result) == 0) {
emit StatusUpdate(obs_module_text(
"AdvSceneSwitcher.condition.video.patternMatchFail"));
} else {

View File

@@ -20,10 +20,12 @@ enum class PreviewType {
class PreviewImage : public QObject {
Q_OBJECT
public:
PreviewImage(std::mutex &);
public slots:
void CreateImage(const VideoInput &, PreviewType,
const PatternMatchParameters &,
const PatternImageData &, ObjDetectParameters,
const PatternMatchParameters &, ObjDetectParameters,
OCRParameters, const AreaParameters &, VideoCondition);
signals:
void ImageReady(const QPixmap &);
@@ -33,6 +35,8 @@ private:
void MarkMatch(QImage &screenshot, const PatternMatchParameters &,
const PatternImageData &, ObjDetectParameters &,
const OCRParameters &, VideoCondition);
std::mutex &_mtx;
};
class PreviewDialog : public QDialog {
@@ -59,9 +63,8 @@ private slots:
signals:
void SelectionAreaChanged(QRect area);
void NeedImage(const VideoInput &, PreviewType,
const PatternMatchParameters &, const PatternImageData &,
ObjDetectParameters, OCRParameters,
const AreaParameters &, VideoCondition);
const PatternMatchParameters &, ObjDetectParameters,
OCRParameters, const AreaParameters &, VideoCondition);
private:
void Start();
@@ -73,7 +76,6 @@ private:
VideoInput _video;
PatternMatchParameters _patternMatchParams;
PatternImageData _patternImageData;
ObjDetectParameters _objDetectParams;
OCRParameters _ocrParams;
AreaParameters _areaParams;