trtexec rejects --minShapes/--optShapes/--maxShapes for a fully static model
("Static model does not take explicit shapes"). TransNetV2's input is fixed at
1x100x27x48x3, so the shape comes from the model itself.
Gallery build now over-fetches TMDB/Wikidata candidates by a configurable
factor: near-duplicate stills (the same photo at different crops or
resolutions) are discarded after embedding, so downloading exactly
images_per_actor left actors short of that many *distinct* embeddings.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
120 lines
3.9 KiB
C++
120 lines
3.9 KiB
C++
#pragma once
|
|
// FaceEmbedderEngine — load SCRFD + ArcFace once, embed many images.
|
|
//
|
|
// Extracted from embed_faces.cpp so the same detect→align→embed pipeline can
|
|
// be driven from a long-lived process (the sae_embed Python module) instead
|
|
// of a fresh CLI invocation per image, which would reload both ONNX sessions
|
|
// every time.
|
|
|
|
#include "config.hpp"
|
|
#include "face_utils.hpp"
|
|
#include "inference/face_detector.hpp"
|
|
#include "inference/face_embedder.hpp"
|
|
#include "types.hpp"
|
|
|
|
#include <opencv2/imgcodecs.hpp>
|
|
#include <opencv2/imgproc.hpp>
|
|
|
|
#include <algorithm>
|
|
#include <array>
|
|
#include <iostream>
|
|
#include <memory>
|
|
#include <string>
|
|
|
|
struct FaceEmbedResult {
|
|
bool ok{false};
|
|
std::string error;
|
|
Embedding embedding{};
|
|
float confidence{0.f};
|
|
float bbox[4]{}; // x, y, w, h
|
|
std::array<cv::Point2f, 5> landmarks{};
|
|
};
|
|
|
|
class FaceEmbedderEngine {
|
|
public:
|
|
// detector_engine/arcface_engine are optional paths to pre-built TensorRT
|
|
// engines. They are required when built with SAE_INFERENCE_BACKEND=TRT
|
|
// (which cannot load .onnx directly) and ignored by the ORT backend.
|
|
FaceEmbedderEngine(const std::string& detector_model,
|
|
const std::string& arcface_model,
|
|
float conf = 0.5f, float nms = 0.4f, int max_side = 500,
|
|
const std::string& detector_engine = "",
|
|
const std::string& arcface_engine = "")
|
|
: max_side_(max_side)
|
|
{
|
|
Config cfg;
|
|
cfg.detector_model = detector_model;
|
|
cfg.arcface_model = arcface_model;
|
|
cfg.detector_engine = detector_engine;
|
|
cfg.arcface_engine = arcface_engine;
|
|
cfg.detector_conf = conf;
|
|
cfg.detector_nms = nms;
|
|
detector_ = make_face_detector(cfg);
|
|
embedder_ = make_face_embedder(cfg);
|
|
}
|
|
|
|
FaceEmbedResult embed_path(const std::string& path) const {
|
|
cv::Mat img = cv::imread(path);
|
|
if (img.empty()) {
|
|
FaceEmbedResult res;
|
|
res.error = "cannot read image";
|
|
return res;
|
|
}
|
|
return embed_mat(img);
|
|
}
|
|
|
|
FaceEmbedResult embed_mat(cv::Mat img) const {
|
|
FaceEmbedResult res;
|
|
|
|
if (max_side_ > 0) {
|
|
const int big = std::max(img.cols, img.rows);
|
|
if (big > max_side_) {
|
|
const double s = static_cast<double>(max_side_) / big;
|
|
cv::resize(img, img, {}, s, s, cv::INTER_AREA);
|
|
}
|
|
}
|
|
|
|
std::vector<DetectedFace> faces = detector_->detect(img);
|
|
if (faces.empty()) {
|
|
cv::Mat enhanced = enhance_for_retry(img);
|
|
faces = detector_->detect(enhanced);
|
|
if (!faces.empty())
|
|
img = enhanced;
|
|
}
|
|
if (faces.empty()) {
|
|
res.error = "no face detected";
|
|
return res;
|
|
}
|
|
if (faces.size() > 1)
|
|
std::cerr << "[warn] " << faces.size()
|
|
<< " faces detected, using highest-confidence one\n";
|
|
|
|
const auto& best = *std::max_element(
|
|
faces.begin(), faces.end(),
|
|
[](const DetectedFace& a, const DetectedFace& b) {
|
|
return a.confidence < b.confidence;
|
|
});
|
|
|
|
cv::Mat crop = align_face(img, best.landmarks);
|
|
if (crop.empty()) {
|
|
res.error = "alignment failed";
|
|
return res;
|
|
}
|
|
|
|
res.ok = true;
|
|
res.embedding = embedder_->embed_one(crop);
|
|
res.confidence = best.confidence;
|
|
res.bbox[0] = best.bbox.x;
|
|
res.bbox[1] = best.bbox.y;
|
|
res.bbox[2] = best.bbox.width;
|
|
res.bbox[3] = best.bbox.height;
|
|
res.landmarks = best.landmarks;
|
|
return res;
|
|
}
|
|
|
|
private:
|
|
std::unique_ptr<IFaceDetector> detector_;
|
|
std::unique_ptr<IFaceEmbedder> embedder_;
|
|
int max_side_;
|
|
};
|