临时提交

This commit is contained in:
dengwenjie
2025-08-29 10:30:35 +08:00
parent 8bf620a330
commit 86ea7eb03e
364 changed files with 8572 additions and 540 deletions

View File

@@ -0,0 +1,58 @@
package cn.smartjavaai.action.config;
import cn.smartjavaai.action.enums.ActionRecModelEnum;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
import lombok.Data;
import java.util.List;
/**
* 人类动作识别模型参数配置
*
* @author dwj
*/
@Data
public class ActionRecModelConfig extends ModelConfig {
/**
* 模型
*/
private ActionRecModelEnum modelEnum;
/**
* 模型路径
*/
private String modelPath;
/**
* 允许的分类列表
*/
private List<String> allowedClasses;
/**
* 置信度阈值
*/
private float threshold = 0.3f;
public ActionRecModelConfig() {
}
public ActionRecModelConfig(ActionRecModelEnum modelEnum, DeviceEnum device) {
this.modelEnum = modelEnum;
setDevice(device);
}
public ActionRecModelConfig(ActionRecModelEnum modelEnum) {
this.modelEnum = modelEnum;
}
}

View File

@@ -0,0 +1,67 @@
package cn.smartjavaai.action.criteria;
import ai.djl.Device;
import ai.djl.modality.Classifications;
import ai.djl.modality.cv.Image;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.action.config.ActionRecModelConfig;
import cn.smartjavaai.action.enums.ActionRecModelEnum;
import cn.smartjavaai.action.model.CommonActionTranslator;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.action.exception.ActionException;
import org.apache.commons.lang3.StringUtils;
import java.nio.file.Paths;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* 动作识别模型工厂
* @author dwj
*/
public class ActionRecCriteriaFactory {
public static Criteria<Image, Classifications> createCriteria(ActionRecModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Criteria<Image, Classifications> criteria = null;
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
params.putAll(config.getCustomParams());
if(config.getModelEnum() == ActionRecModelEnum.VIT_BASE_PATCH16_224){
criteria =
Criteria.builder()
.setTypes(Image.class, Classifications.class)
.optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null :
config.getModelEnum().getModelUri())
.optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null)
.optEngine("PyTorch")
.optDevice(device)
.optProgress(new ProgressBar())
.build();
}else {
if (StringUtils.isBlank(config.getModelPath())){
throw new ActionException("请指定模型路径");
}
int width = 224;
int height = 224;
if(config.getModelEnum() == ActionRecModelEnum.INCEPTIONV3_KINETICS400){
width = 299;
height = 299;
}
criteria =
Criteria.builder()
.setTypes(Image.class, Classifications.class)
.optTranslator(new CommonActionTranslator(width, height))
.optEngine("OnnxRuntime")
.optModelPath(Paths.get(config.getModelPath()))
.optDevice(device)
.optProgress(new ProgressBar())
.build();
}
return criteria;
}
}

View File

@@ -0,0 +1,48 @@
package cn.smartjavaai.action.enums;
/**
* 人类动作识别模型枚举
* @author dwj
*/
public enum ActionRecModelEnum {
VIT_BASE_PATCH16_224("djl://ai.djl.pytorch/Human-Action-Recognition-VIT-Base-patch16-224"),
INCEPTIONV3_KINETICS400(""),
INCEPTIONV1_KINETICS400(""),
RESNET18_V1B_KINETICS400(""),
RESNET34_V1B_KINETICS400(""),
RESNET50_V1B_KINETICS400(""),
RESNET101_V1B_KINETICS400(""),
RESNET152_V1B_KINETICS400("");
/**
* 根据名称获取枚举 (忽略大小写和下划线变体)
*/
public static ActionRecModelEnum fromName(String name) {
String formatted = name.trim().toUpperCase().replaceAll("[-_]", "");
for (ActionRecModelEnum model : values()) {
if (model.name().replaceAll("_", "").equals(formatted)) {
return model;
}
}
throw new IllegalArgumentException("未知模型名称: " + name);
}
private final String modelUri;
ActionRecModelEnum(String modelUri) {
this.modelUri = modelUri;
}
public String getModelUri() {
return modelUri;
}
}

View File

@@ -0,0 +1,29 @@
package cn.smartjavaai.action.exception;
/**
* 动作检测异常
* @author dwj
*/
public class ActionException extends RuntimeException{
public ActionException() {
super();
}
public ActionException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) {
super(message, cause, enableSuppression, writableStackTrace);
}
public ActionException(String message, Throwable cause) {
super(message, cause);
}
public ActionException(String message) {
super(message);
}
public ActionException(Throwable cause) {
super(cause);
}
}

View File

@@ -0,0 +1,46 @@
package cn.smartjavaai.action.model;
import ai.djl.modality.Classifications;
import ai.djl.modality.cv.Image;
import cn.smartjavaai.action.config.ActionRecModelConfig;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
import java.awt.image.BufferedImage;
import java.io.InputStream;
/**
* 人类动作识别模型
* @author dwj
* @date 2025/8/12
*/
public interface ActionRecModel extends AutoCloseable{
/**
* 加载模型
* @param config
*/
void loadModel(ActionRecModelConfig config);
/**
* 动作检测
* @param base64Image
* @return
*/
default R<Classifications> detectBase64(String base64Image){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 动作检测
* @param image
* @return
*/
default R<Classifications> detect(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
}

View File

@@ -0,0 +1,142 @@
package cn.smartjavaai.action.model;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.Classifications;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import cn.hutool.core.img.ImgUtil;
import cn.smartjavaai.action.config.ActionRecModelConfig;
import cn.smartjavaai.action.criteria.ActionRecCriteriaFactory;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.Base64ImageUtils;
import cn.smartjavaai.objectdetection.criteria.CriteriaBuilderFactory;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.vision.utils.ClassificationFilter;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.impl.GenericObjectPool;
import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.util.Collections;
import java.util.Objects;
/**
* 通用人类动作识别模型
* @author dwj
*/
@Slf4j
public class CommonActionRecModel implements ActionRecModel{
private ActionRecModelConfig config;
private ZooModel<Image, Classifications> model;
private GenericObjectPool<Predictor<Image, Classifications>> predictorPool;
@Override
public void loadModel(ActionRecModelConfig config) {
if(Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型枚举");
}
Criteria<Image, Classifications> criteria = ActionRecCriteriaFactory.createCriteria(config);
this.config = config;
try {
model = criteria.loadModel();
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model));
int predictorPoolSize = config.getPredictorPoolSize();
if(config.getPredictorPoolSize() <= 0){
predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数
}
predictorPool.setMaxTotal(predictorPoolSize);
log.debug("当前设备: " + model.getNDManager().getDevice());
log.debug("当前引擎: " + Engine.getInstance().getEngineName());
log.debug("模型推理器线程池最大数量: " + predictorPoolSize);
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new DetectionException("模型加载失败", e);
}
}
@Override
public R<Classifications> detectBase64(String base64Image) {
if(StringUtils.isBlank(base64Image)){
return R.fail(R.Status.INVALID_IMAGE);
}
try {
byte[] imageData = Base64ImageUtils.base64ToImage(base64Image);
Image image = ImageFactory.getInstance().fromInputStream(new ByteArrayInputStream(imageData));
return detect(image);
} catch (IOException e) {
throw new DetectionException("读取图片异常", e);
}
}
@Override
public R<Classifications> detect(Image image) {
Classifications classifications = detectCore(image);
// 过滤
if(config.getThreshold() > 0 && CollectionUtils.isNotEmpty(config.getAllowedClasses())
&& Objects.nonNull(classifications) && !classifications.items().isEmpty()){
classifications = new ClassificationFilter(config.getAllowedClasses(), config.getThreshold()).filter(classifications);
}
return R.ok(classifications);
}
/**
* 模型核心推理方法
* @param image
* @return
*/
public Classifications detectCore(Image image) {
Predictor<Image, Classifications> predictor = null;
try {
predictor = predictorPool.borrowObject();
return predictor.predict(image);
} catch (Exception e) {
throw new DetectionException("动作识别错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
log.debug("释放资源");
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public void close() throws Exception {
try {
if (predictorPool != null) {
predictorPool.close();
}
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (model != null) {
model.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,140 @@
package cn.smartjavaai.action.model;
import ai.djl.modality.Classifications;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.util.NDImageUtils;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import ai.djl.translate.Batchifier;
import ai.djl.translate.Translator;
import ai.djl.translate.TranslatorContext;
import cn.smartjavaai.common.utils.LetterBoxUtils;
import java.util.Arrays;
/**
* @author dwj
*/
public class CommonActionTranslator implements Translator<Image, Classifications> {
private int width = 224;
private int height = 224;
private String[] labels;
public static String rawLabels = "'abseiling','air_drumming','answering_questions','applauding','applying_cream','archery'," +
"'arm_wrestling','arranging_flowers','assembling_computer','auctioning','baby_waking_up','baking_cookies'," +
"'balloon_blowing','bandaging','barbequing','bartending','beatboxing','bee_keeping','belly_dancing'," +
"'bench_pressing','bending_back','bending_metal','biking_through_snow','blasting_sand','blowing_glass'," +
"'blowing_leaves','blowing_nose','blowing_out_candles','bobsledding','bookbinding','bouncing_on_trampoline'," +
"'bowling','braiding_hair','breading_or_breadcrumbing','breakdancing','brush_painting','brushing_hair'," +
"'brushing_teeth','building_cabinet','building_shed','bungee_jumping','busking','canoeing_or_kayaking'," +
"'capoeira','carrying_baby','cartwheeling','carving_pumpkin','catching_fish','catching_or_throwing_baseball'," +
"'catching_or_throwing_frisbee','catching_or_throwing_softball','celebrating','changing_oil','changing_wheel'," +
"'checking_tires','cheerleading','chopping_wood','clapping','clay_pottery_making','clean_and_jerk'," +
"'cleaning_floor','cleaning_gutters','cleaning_pool','cleaning_shoes','cleaning_toilet','cleaning_windows'," +
"'climbing_a_rope','climbing_ladder','climbing_tree','contact_juggling','cooking_chicken','cooking_egg'," +
"'cooking_on_campfire','cooking_sausages','counting_money','country_line_dancing','cracking_neck','crawling_baby'," +
"'crossing_river','crying','curling_hair','cutting_nails','cutting_pineapple','cutting_watermelon'," +
"'dancing_ballet','dancing_charleston','dancing_gangnam_style','dancing_macarena','deadlifting'," +
"'decorating_the_christmas_tree','digging','dining','disc_golfing','diving_cliff','dodgeball','doing_aerobics'," +
"'doing_laundry','doing_nails','drawing','dribbling_basketball','drinking','drinking_beer','drinking_shots'," +
"'driving_car','driving_tractor','drop_kicking','drumming_fingers','dunking_basketball','dying_hair'," +
"'eating_burger','eating_cake','eating_carrots','eating_chips','eating_doughnuts','eating_hotdog'," +
"'eating_ice_cream','eating_spaghetti','eating_watermelon','egg_hunting','exercising_arm'," +
"'exercising_with_an_exercise_ball','extinguishing_fire','faceplanting','feeding_birds','feeding_fish'," +
"'feeding_goats','filling_eyebrows','finger_snapping','fixing_hair','flipping_pancake','flying_kite'," +
"'folding_clothes','folding_napkins','folding_paper','front_raises','frying_vegetables','garbage_collecting'," +
"'gargling','getting_a_haircut','getting_a_tattoo','giving_or_receiving_award','golf_chipping','golf_driving'," +
"'golf_putting','grinding_meat','grooming_dog','grooming_horse','gymnastics_tumbling','hammer_throw'," +
"'headbanging','headbutting','high_jump','high_kick','hitting_baseball','hockey_stop','holding_snake'," +
"'hopscotch','hoverboarding','hugging','hula_hooping','hurdling','hurling_-sport-','ice_climbing','ice_fishing'," +
"'ice_skating','ironing','javelin_throw','jetskiing','jogging','juggling_balls','juggling_fire'," +
"'juggling_soccer_ball','jumping_into_pool','jumpstyle_dancing','kicking_field_goal','kicking_soccer_ball'," +
"'kissing','kitesurfing','knitting','krumping','laughing','laying_bricks','long_jump','lunge','making_a_cake'," +
"'making_a_sandwich','making_bed','making_jewelry','making_pizza','making_snowman','making_sushi','making_tea'," +
"'marching','massaging_back','massaging_feet','massaging_legs','massaging_person's_head','milking_cow'," +
"'mopping_floor','motorcycling','moving_furniture','mowing_lawn','news_anchoring','opening_bottle'," +
"'opening_present','paragliding','parasailing','parkour','passing_American_football_-in_game-'," +
"'passing_American_football_-not_in_game-','peeling_apples','peeling_potatoes','petting_animal_-not_cat-'," +
"'petting_cat','picking_fruit','planting_trees','plastering','playing_accordion','playing_badminton'," +
"'playing_bagpipes','playing_basketball','playing_bass_guitar','playing_cards','playing_cello','playing_chess'," +
"'playing_clarinet','playing_controller','playing_cricket','playing_cymbals','playing_didgeridoo','playing_drums'," +
"'playing_flute','playing_guitar','playing_harmonica','playing_harp','playing_ice_hockey','playing_keyboard'," +
"'playing_kickball','playing_monopoly','playing_organ','playing_paintball','playing_piano','playing_poker'," +
"'playing_recorder','playing_saxophone','playing_squash_or_racquetball','playing_tennis','playing_trombone'," +
"'playing_trumpet','playing_ukulele','playing_violin','playing_volleyball','playing_xylophone','pole_vault'," +
"'presenting_weather_forecast','pull_ups','pumping_fist','pumping_gas','punching_bag','punching_person_-boxing-'," +
"'push_up','pushing_car','pushing_cart','pushing_wheelchair','reading_book','reading_newspaper','recording_music'," +
"'riding_a_bike','riding_camel','riding_elephant','riding_mechanical_bull','riding_mountain_bike','riding_mule'," +
"'riding_or_walking_with_horse','riding_scooter','riding_unicycle','ripping_paper','robot_dancing','rock_climbing'," +
"'rock_scissors_paper','roller_skating','running_on_treadmill','sailing','salsa_dancing','sanding_floor'," +
"'scrambling_eggs','scuba_diving','setting_table','shaking_hands','shaking_head','sharpening_knives'," +
"'sharpening_pencil','shaving_head','shaving_legs','shearing_sheep','shining_shoes','shooting_basketball'," +
"'shooting_goal_-soccer-','shot_put','shoveling_snow','shredding_paper','shuffling_cards','side_kick'," +
"'sign_language_interpreting','singing','situp','skateboarding','ski_jumping','skiing_-not_slalom_or_crosscountry-'," +
"'skiing_crosscountry','skiing_slalom','skipping_rope','skydiving','slacklining','slapping','sled_dog_racing'," +
"'smoking','smoking_hookah','snatch_weight_lifting','sneezing','sniffing','snorkeling','snowboarding','snowkiting'," +
"'snowmobiling','somersaulting','spinning_poi','spray_painting','spraying','springboard_diving','squat'," +
"'sticking_tongue_out','stomping_grapes','stretching_arm','stretching_leg','strumming_guitar','surfing_crowd'," +
"'surfing_water','sweeping_floor','swimming_backstroke','swimming_breast_stroke','swimming_butterfly_stroke'," +
"'swing_dancing','swinging_legs','swinging_on_something','sword_fighting','tai_chi','taking_a_shower','tango_dancing'," +
"'tap_dancing','tapping_guitar','tapping_pen','tasting_beer','tasting_food','testifying','texting','throwing_axe'," +
"'throwing_ball','throwing_discus','tickling','tobogganing','tossing_coin','tossing_salad','training_dog'," +
"'trapezing','trimming_or_shaving_beard','trimming_trees','triple_jump','tying_bow_tie','tying_knot_-not_on_a_tie-'," +
"'tying_tie','unboxing','unloading_truck','using_computer','using_remote_controller_-not_gaming-','using_segway'," +
"'vault','waiting_in_line','walking_the_dog','washing_dishes','washing_feet','washing_hair','washing_hands'," +
"'water_skiing','water_sliding','watering_plants','waxing_back','waxing_chest','waxing_eyebrows','waxing_legs'," +
"'weaving_basket','welding','whistling','windsurfing','wrapping_present','wrestling','writing','yawning','yoga','zumba'";
public CommonActionTranslator(int width, int height) {
this.width = width;
this.height = height;
}
@Override
public void prepare(TranslatorContext ctx) throws Exception {
labels = rawLabels.replace("'", "").split(",");
labels = Arrays.stream(labels)
.map(String::trim)
.toArray(String[]::new);
}
@Override
public NDList processInput(TranslatorContext ctx, Image input) {
NDManager manager = ctx.getNDManager();
NDArray array = input.toNDArray(manager, Image.Flag.COLOR); // HWC, uint8
array = NDImageUtils.resize(array, width, height);
// 转 float32
array = array.toType(DataType.FLOAT32, false);
// 自己归一化注意均值和标准差要乘以255因为原始是0~255
float[] mean = {0.485f * 255, 0.456f * 255, 0.406f * 255};
float[] std = {0.229f * 255, 0.224f * 255, 0.225f * 255};
// 增加 batch 维度,变成 (1, H, W, C)
array = array.expandDims(0);
System.out.println(Arrays.toString(array.getShape().getShape()));
return new NDList(array);
}
@Override
public Classifications processOutput(TranslatorContext ctx, NDList list) {
NDArray output = list.singletonOrThrow();
output = output.softmax(1); // 计算概率
Classifications classifications = new Classifications(Arrays.asList(labels), output);
return classifications;
}
@Override
public Batchifier getBatchifier() {
return null;
}
}

View File

@@ -0,0 +1,55 @@
package cn.smartjavaai.instanceseg.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.instanceseg.enums.InstanceSegModelEnum;
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
import lombok.Data;
import java.util.List;
/**
* 实例分割模型参数配置
*
* @author dwj
*/
@Data
public class InstanceSegModelConfig extends ModelConfig {
/**
* 模型
*/
private InstanceSegModelEnum modelEnum;
/**
* 模型路径
*/
private String modelPath;
/**
* 允许的分类列表
*/
private List<String> allowedClasses;
/**
* 置信度阈值
*/
private float threshold = 0.3f;
public InstanceSegModelConfig() {
}
public InstanceSegModelConfig(InstanceSegModelEnum modelEnum, DeviceEnum device) {
this.modelEnum = modelEnum;
setDevice(device);
}
public InstanceSegModelConfig(InstanceSegModelEnum modelEnum) {
this.modelEnum = modelEnum;
}
}

View File

@@ -0,0 +1,47 @@
package cn.smartjavaai.instanceseg.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig;
import cn.smartjavaai.instanceseg.translator.YoloSegmentationTranslatorFactory2;
import org.apache.commons.lang3.StringUtils;
import java.nio.file.Paths;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* 实例分割Criteria工厂
* @author dwj
*/
public class InstanceSegCriteriaFactory {
public static Criteria<Image, DetectedObjects> createCriteria(InstanceSegModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Criteria<Image, DetectedObjects> criteria = null;
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
params.putAll(config.getCustomParams());
// YoloV5Translator.Builder builder = new YoloV5Translator.Builder()
// .optSynsetArtifactName("synset.txt").setPipeline()
criteria =
Criteria.builder()
.setTypes(Image.class, DetectedObjects.class)
.optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null :
config.getModelEnum().getModelUri())
.optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null)
.optDevice(device)
.optEngine("PyTorch")
.optTranslatorFactory(new YoloSegmentationTranslatorFactory2())
.optProgress(new ProgressBar())
.build();
return criteria;
}
}

View File

@@ -0,0 +1,17 @@
package cn.smartjavaai.instanceseg.entity;
import lombok.Data;
/**
* 检测参数
* @author dwj
*/
@Data
public class DetectParams {
/**
* 置信度阈值
*/
private float threshold = 0.3f;
}

View File

@@ -0,0 +1,43 @@
package cn.smartjavaai.instanceseg.enums;
/**
* 实例分割模型枚举
* @author dwj
*/
public enum InstanceSegModelEnum {
SEG_YOLO11N_PYTORCH("djl://ai.djl.pytorch/yolo11n-seg"),
SEG_YOLOV8N_PYTORCH("djl://ai.djl.pytorch/yolo11n-seg"),
SEG_YOLO11N_ONNX("djl://ai.djl.onnxruntime/yolo11n-seg"),
SEG_YOLOV8N_ONNX("djl://ai.djl.onnxruntime/yolov8n-seg"),
SEG_MASK_RCNN("djl://ai.djl.mxnet/mask_rcnn");
/**
* 根据名称获取枚举 (忽略大小写和下划线变体)
*/
public static InstanceSegModelEnum fromName(String name) {
String formatted = name.trim().toUpperCase().replaceAll("[-_]", "");
for (InstanceSegModelEnum model : values()) {
if (model.name().replaceAll("_", "").equals(formatted)) {
return model;
}
}
throw new IllegalArgumentException("未知模型名称: " + name);
}
private final String modelUri;
InstanceSegModelEnum(String modelUri) {
this.modelUri = modelUri;
}
public String getModelUri() {
return modelUri;
}
}

View File

@@ -0,0 +1,29 @@
package cn.smartjavaai.instanceseg.exception;
/**
* 实例分割异常
* @author dwj
*/
public class InstanceSegException extends RuntimeException{
public InstanceSegException() {
super();
}
public InstanceSegException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) {
super(message, cause, enableSuppression, writableStackTrace);
}
public InstanceSegException(String message, Throwable cause) {
super(message, cause);
}
public InstanceSegException(String message) {
super(message);
}
public InstanceSegException(Throwable cause) {
super(cause);
}
}

View File

@@ -0,0 +1,153 @@
package cn.smartjavaai.instanceseg.model;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import cn.smartjavaai.common.cv.SmartImageFactory;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig;
import cn.smartjavaai.instanceseg.criteria.InstanceSegCriteriaFactory;
import cn.smartjavaai.instanceseg.exception.InstanceSegException;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.vision.utils.DetectedObjectsFilter;
import cn.smartjavaai.vision.utils.DetectorUtils;
import cn.smartjavaai.vision.utils.CategoryMaskFilter;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.pool2.impl.GenericObjectPool;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.Objects;
/**
* 实例分割模型
* @author dwj
*/
@Slf4j
public class CommonInstanceSegModel implements InstanceSegModel {
private InstanceSegModelConfig config;
private ZooModel<Image, DetectedObjects> model;
private GenericObjectPool<Predictor<Image, DetectedObjects>> predictorPool;
@Override
public void loadModel(InstanceSegModelConfig config) {
if(Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型枚举");
}
Criteria<Image, DetectedObjects> criteria = InstanceSegCriteriaFactory.createCriteria(config);
this.config = config;
try {
model = criteria.loadModel();
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model));
int predictorPoolSize = config.getPredictorPoolSize();
if(config.getPredictorPoolSize() <= 0){
predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数
}
predictorPool.setMaxTotal(predictorPoolSize);
log.debug("当前设备: " + model.getNDManager().getDevice());
log.debug("当前引擎: " + Engine.getInstance().getEngineName());
log.debug("模型推理器线程池最大数量: " + predictorPoolSize);
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new DetectionException("模型加载失败", e);
}
}
@Override
public R<DetectionResponse> detect(Image image) {
DetectedObjects detectedObjects = detectCore(image);
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image);
return R.ok(detectionResponse);
}
/**
* 模型核心推理方法
* @param image
* @return
*/
@Override
public DetectedObjects detectCore(Image image) {
Predictor<Image, DetectedObjects> predictor = null;
try {
predictor = predictorPool.borrowObject();
DetectedObjects detectedObjects = predictor.predict(image);
//过滤
if(Objects.nonNull(detectedObjects) && detectedObjects.getNumberOfObjects() > 0){
DetectedObjectsFilter detectedObjectsFilter = new DetectedObjectsFilter(config.getAllowedClasses(), config.getThreshold());
detectedObjects = detectedObjectsFilter.filter(detectedObjects);
}
return detectedObjects;
} catch (Exception e) {
throw new DetectionException("实例分割错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
log.debug("释放资源");
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public R<DetectionResponse> detectAndDraw(Image image) {
DetectedObjects detectedObjects = detectCore(image);
image.drawBoundingBoxes(detectedObjects);
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image);
detectionResponse.setDrawnImage(image);
return R.ok(detectionResponse);
}
@Override
public R<DetectionResponse> detectAndDraw(String imagePath, String outputPath) {
try {
Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath));
DetectedObjects detectedObjects = detectCore(img);
img.drawBoundingBoxes(detectedObjects);
img.save(Files.newOutputStream(Paths.get(outputPath)), "png");
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, img);
return R.ok(detectionResponse);
} catch (IOException e) {
throw new InstanceSegException(e);
}
}
@Override
public void close() throws Exception {
try {
if (predictorPool != null) {
predictorPool.close();
}
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (model != null) {
model.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,47 @@
package cn.smartjavaai.instanceseg.model;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.CategoryMask;
import ai.djl.modality.cv.output.DetectedObjects;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig;
import java.awt.image.BufferedImage;
/**
* 实例分割模型
* @author dwj
*/
public interface InstanceSegModel extends AutoCloseable{
/**
* 加载模型
* @param config
*/
void loadModel(InstanceSegModelConfig config);
/**
* 实例分割
* @param image
* @return
*/
default R<DetectionResponse> detect(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
default DetectedObjects detectCore(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
default R<DetectionResponse> detectAndDraw(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
default R<DetectionResponse> detectAndDraw(String imagePath, String outputPath){
throw new UnsupportedOperationException("默认不支持该功能");
}
}

View File

@@ -0,0 +1,182 @@
/*
* Copyright 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved.
*
* Licensed under the Apache License, Version 2.0 (the "License"). You may not use this file except in compliance
* with the License. A copy of the License is located at
*
* http://aws.amazon.com/apache2.0/
*
* or in the "license" file accompanying this file. This file is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES
* OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions
* and limitations under the License.
*/
package cn.smartjavaai.instanceseg.translator;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.output.Mask;
import ai.djl.modality.cv.output.Rectangle;
import ai.djl.modality.cv.transform.Resize;
import ai.djl.modality.cv.transform.ToTensor;
import ai.djl.modality.cv.translator.YoloV5Translator;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.types.DataType;
import ai.djl.translate.ArgumentsUtil;
import ai.djl.translate.Pipeline;
import ai.djl.translate.TranslatorContext;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
/** A translator for Yolov8 instance segmentation models. */
public class YoloSegmentationTranslator2 extends YoloV5Translator {
private static final int[] AXIS_0 = {0};
private static final int[] AXIS_1 = {1};
private float threshold;
private float nmsThreshold;
/**
* Creates the instance segmentation translator from the given builder.
*
* @param builder the builder for the translator
*/
public YoloSegmentationTranslator2(Builder builder) {
super(builder);
this.threshold = 0.25f;
this.nmsThreshold = 0.4F;
}
/** {@inheritDoc} */
@Override
public DetectedObjects processOutput(TranslatorContext ctx, NDList list) {
NDArray pred = list.get(0);
NDArray protos = list.get(1);
int maskIndex = classes.size() + 4;
NDArray candidates = pred.get("4:" + maskIndex).max(AXIS_0).gt(threshold);
pred = pred.transpose();
NDArray sub = pred.get("..., :4");
sub = xywh2xyxy(sub);
pred = sub.concat(pred.get("..., 4:"), -1);
pred = pred.get(candidates);
NDList split = pred.split(new long[] {4, maskIndex}, 1);
NDArray box = split.get(0);
int numBox = Math.toIntExact(box.getShape().get(0));
float[] buf = box.toFloatArray();
float[] confidences = split.get(1).max(AXIS_1).toFloatArray();
long[] ids = split.get(1).argMax(1).toLongArray();
List<Rectangle> boxes = new ArrayList<>(numBox);
List<Double> scores = new ArrayList<>(numBox);
for (int i = 0; i < numBox; ++i) {
float xPos = buf[i * 4];
float yPos = buf[i * 4 + 1];
float w = buf[i * 4 + 2] - xPos;
float h = buf[i * 4 + 3] - yPos;
Rectangle rect = new Rectangle(xPos, yPos, w, h);
boxes.add(rect);
scores.add((double) confidences[i]);
}
List<Integer> nms = Rectangle.nms(boxes, scores, nmsThreshold);
long[] idx = nms.stream().mapToLong(Integer::longValue).toArray();
NDArray selected = box.getManager().create(idx);
NDArray masks = split.get(2).get(selected);
int maskW = Math.toIntExact(protos.getShape().get(2));
int maskH = Math.toIntExact(protos.getShape().get(1));
protos = protos.reshape(32, (long) maskH * maskW);
masks =
masks.matMul(protos)
.reshape(nms.size(), maskH, maskW)
.gt(0f)
.toType(DataType.FLOAT32, true);
float[] maskArray = masks.toFloatArray();
box = box.get(selected);
buf = box.toFloatArray();
List<String> retClasses = new ArrayList<>();
List<Double> retProbs = new ArrayList<>();
List<BoundingBox> retBB = new ArrayList<>();
for (int i = 0; i < idx.length; ++i) {
float x = buf[i * 4] / width;
float y = buf[i * 4 + 1] / height;
float w = buf[i * 4 + 2] / width - x;
float h = buf[i * 4 + 3] / width - y;
int id = nms.get(i);
retClasses.add(classes.get((int) ids[id]));
retProbs.add((double) confidences[id]);
float[][] maskFloat = new float[maskH][maskW];
int pos = i * maskH * maskW;
for (int j = 0; j < maskH; j++) {
System.arraycopy(maskArray, pos + j * maskW, maskFloat[j], 0, maskW);
}
Mask bb = new Mask(x, y, w, h, maskFloat, true);
retBB.add(bb);
}
return new DetectedObjects(retClasses, retProbs, retBB);
}
private NDArray xywh2xyxy(NDArray array) {
NDArray xy = array.get("..., :2", new Object[0]);
NDArray wh = array.get("..., 2:", new Object[0]).div(2);
return xy.sub(wh).concat(xy.add(wh), -1);
}
/**
* Creates a builder to build a {@code YoloSegmentationTranslator}.
*
* @return a new builder
*/
public static Builder builder() {
return new Builder();
}
/**
* Creates a builder to build a {@code YoloSegmentationTranslator} with specified arguments.
*
* @param arguments arguments to specify builder options
* @return a new builder
*/
public static Builder builder(Map<String, ?> arguments) {
Builder builder = new Builder();
builder.optSynsetArtifactName("synset.txt");
builder.setImageSize(640, 640);
return builder;
}
/** The builder for instance segmentation translator. */
public static class Builder extends YoloV5Translator.Builder {
Builder() {}
/** {@inheritDoc} */
@Override
protected Builder self() {
return this;
}
/** {@inheritDoc} */
@Override
public YoloSegmentationTranslator2 build() {
pipeline = new Pipeline();
pipeline.add(new Resize(640, 640));
pipeline.add(new ToTensor());
// validate();
return new YoloSegmentationTranslator2(this);
}
}
}

View File

@@ -0,0 +1,37 @@
/*
* Copyright 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved.
*
* Licensed under the Apache License, Version 2.0 (the "License"). You may not use this file except in compliance
* with the License. A copy of the License is located at
*
* http://aws.amazon.com/apache2.0/
*
* or in the "license" file accompanying this file. This file is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES
* OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions
* and limitations under the License.
*/
package cn.smartjavaai.instanceseg.translator;
import ai.djl.Model;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.translator.ObjectDetectionTranslatorFactory;
import ai.djl.translate.Translator;
import java.io.Serializable;
import java.util.Map;
/** A translatorFactory that creates a {@link ai.djl.modality.cv.translator.YoloSegmentationTranslator} instance. */
public class YoloSegmentationTranslatorFactory2 extends ObjectDetectionTranslatorFactory
implements Serializable {
private static final long serialVersionUID = 1L;
/** {@inheritDoc} */
@Override
protected Translator<Image, DetectedObjects> buildBaseTranslator(
Model model, Map<String, ?> arguments) {
Translator<Image, DetectedObjects> translator = YoloSegmentationTranslator2.builder(arguments).build();
return translator;
}
}

View File

@@ -0,0 +1,69 @@
package cn.smartjavaai.objectdetection.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
import lombok.Data;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
/**
* 目标检测模型参数配置
*
* @author dwj
* @date 2025/4/4
*/
@Data
public class DetectorModelConfig extends ModelConfig {
/**
* 模型
*/
private DetectorModelEnum modelEnum;
/**
* 置信度阈值
*/
private float threshold = DetectorConstant.DEFAULT_THRESHOLD;
/**
* 模型路径
*/
private String modelPath;
/**
* 候选框数量默认为8400. 应设置0到8400之间的整数
* 用于性能优化的关键参数它通过限制模型后处理阶段需要处理的候选框bounding boxes数量来提高推理速度
* 建议不低于1000
*/
private int maxBox;
/**
* 允许的分类列表
*/
private List<String> allowedClasses;
/**
* 检测结果数量
*/
private int topK;
public DetectorModelConfig() {
}
public DetectorModelConfig(DetectorModelEnum modelEnum, DeviceEnum device) {
this.modelEnum = modelEnum;
setDevice(device);
}
public DetectorModelConfig(DetectorModelEnum modelEnum) {
this.modelEnum = modelEnum;
}
}

View File

@@ -0,0 +1,62 @@
package cn.smartjavaai.objectdetection.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
import cn.smartjavaai.objectdetection.enums.PersonDetectorModelEnum;
import lombok.Data;
import java.util.List;
/**
* 行人检测模型参数配置
*
* @author dwj
* @date 2025/4/4
*/
@Data
public class PersonDetModelConfig extends ModelConfig {
/**
* 模型
*/
private PersonDetectorModelEnum modelEnum;
/**
* 置信度阈值
*/
private float threshold;
/**
* 模型路径
*/
private String modelPath;
/**
* 允许的分类列表
*/
private List<String> allowedClasses;
/**
* 按置信度分数排序后,最多保留的检测框数量
*/
private int topK;
public PersonDetModelConfig() {
}
public PersonDetModelConfig(PersonDetectorModelEnum modelEnum, DeviceEnum device) {
this.modelEnum = modelEnum;
setDevice(device);
}
public PersonDetModelConfig(PersonDetectorModelEnum modelEnum) {
this.modelEnum = modelEnum;
}
}

View File

@@ -0,0 +1,17 @@
package cn.smartjavaai.objectdetection.constant;
/**
* @author dwj
* @date 2025/4/7
*/
public class DetectorConstant {
/**
* 置信度阈值
*/
public static final float DEFAULT_THRESHOLD = 0.5F;
}

View File

@@ -0,0 +1,46 @@
package cn.smartjavaai.objectdetection.criteria;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import org.apache.commons.lang3.StringUtils;
/**
* Criteria构建工厂
* @author dwj
* @date 2025/5/14
*/
public class CriteriaBuilderFactory {
public static Criteria<Image, DetectedObjects> createCriteria(DetectorModelConfig config) {
//以下模型modelPath不允许为空
if(config.getModelEnum() == DetectorModelEnum.YOLOV8_OFFICIAL ||
config.getModelEnum() == DetectorModelEnum.YOLOV12_OFFICIAL ||
config.getModelEnum() == DetectorModelEnum.YOLOV8_CUSTOM ||
config.getModelEnum() == DetectorModelEnum.YOLOV12_CUSTOM ||
config.getModelEnum() == DetectorModelEnum.TENSORFLOW2_OFFICIAL){
if(StringUtils.isBlank(config.getModelPath())){
throw new DetectionException("modelPath is null");
}
}
switch (config.getModelEnum()) {
case YOLOV8_OFFICIAL:
return new YoloCriteriaBuilder().buildCriteria(config);
case YOLOV12_OFFICIAL:
return new YoloCriteriaBuilder().buildCriteria(config);
case YOLOV8_CUSTOM:
return new YoloCriteriaBuilder().buildCriteria(config);
case YOLOV12_CUSTOM:
return new YoloCriteriaBuilder().buildCriteria(config);
case TENSORFLOW2_OFFICIAL:
return new Tensorflow2CriteriaBuilder().buildCriteria(config);
// 其他类型
default:
return new DJLModelCriteriaBuilder().buildCriteria(config);
}
}
}

View File

@@ -0,0 +1,22 @@
package cn.smartjavaai.objectdetection.criteria;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
/**
* 模型加载策略接口,用于根据不同模型类型构建对应的 DJL Criteria 实例
* @author dwj
* @date 2025/5/14
*/
public interface CriteriaBuilderStrategy {
/**
* 根据模型类型构建对应的 DJL Criteria 实例
* @param config
* @return
*/
Criteria<Image, DetectedObjects> buildCriteria(DetectorModelConfig config);
}

View File

@@ -0,0 +1,41 @@
package cn.smartjavaai.objectdetection.criteria;
import ai.djl.Application;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
import java.util.Objects;
/**
* DJL提供的Criteria 构建器
* @author dwj
* @date 2025/5/14
*/
public class DJLModelCriteriaBuilder implements CriteriaBuilderStrategy {
private static final String DJL_MODEL_PREFIX = "djl://";
@Override
public Criteria<Image, DetectedObjects> buildCriteria(DetectorModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Criteria<Image, DetectedObjects> criteria = Criteria.builder()
.optApplication(Application.CV.OBJECT_DETECTION)
.setTypes(Image.class, DetectedObjects.class)
.optArgument("threshold", config.getThreshold() > 0 ? config.getThreshold() : DetectorConstant.DEFAULT_THRESHOLD)
.optModelUrls(DJL_MODEL_PREFIX + config.getModelEnum().getModelUri())
.optDevice(device)
//.optOption("ortDevice", "TensorRT")
.optProgress(new ProgressBar())
.build();
return criteria;
}
}

View File

@@ -0,0 +1,76 @@
package cn.smartjavaai.objectdetection.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import ai.djl.translate.Translator;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.obb.config.ObbDetModelConfig;
import cn.smartjavaai.obb.entity.ObbResult;
import cn.smartjavaai.obb.enums.ObbDetModelEnum;
import cn.smartjavaai.obb.exception.ObbDetException;
import cn.smartjavaai.obb.translator.YoloV11OddTranslator;
import cn.smartjavaai.objectdetection.config.PersonDetModelConfig;
import cn.smartjavaai.objectdetection.enums.PersonDetectorModelEnum;
import cn.smartjavaai.objectdetection.translator.YoloV8PersonDetTranslator;
import org.apache.commons.lang3.StringUtils;
import java.nio.file.Paths;
import java.util.HashMap;
import java.util.Map;
import java.util.Objects;
/**
* 行人检测Criteria创建工厂
* @author dwj
*/
public class PersonDetCriteriaFactory {
/**
* 创建行人检测Criteria
* @param config
* @return
*/
public static Criteria<Image, DetectedObjects> createCriteria(PersonDetModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Translator<Image, DetectedObjects> translator = getTranslator(config);
//检查模型路径
if (StringUtils.isBlank(config.getModelPath())){
throw new ObbDetException("请指定模型路径");
}
Criteria<Image, DetectedObjects> criteria =
Criteria.builder()
.setTypes(Image.class, DetectedObjects.class)
.optModelPath(Paths.get(config.getModelPath()))
.optTranslator(translator)
.optDevice(device)
.optProgress(new ProgressBar())
.optEngine(config.getModelEnum().getEngine())
.build();
return criteria;
}
/**
* 获取行人检测Translator
* @param config
* @return
*/
public static Translator<Image, DetectedObjects> getTranslator(PersonDetModelConfig config) {
Translator<Image, DetectedObjects> translator = null;
if (config.getModelEnum() == PersonDetectorModelEnum.YOLOV8_PERSON){
translator = YoloV8PersonDetTranslator.builder()
.setImageSize(config.getModelEnum().getInputSize(), config.getModelEnum().getInputSize())
.optThreshold(config.getThreshold() > 0 ? config.getThreshold() : 0.5f)
.optNmsThreshold(0.45f).build();
}
return translator;
}
}

View File

@@ -0,0 +1,119 @@
package cn.smartjavaai.objectdetection.criteria;
import ai.djl.Application;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.translator.YoloV8TranslatorFactory;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.objectdetection.translator.TensorflowTranslator;
import cn.smartjavaai.vision.utils.TensorflowSynsetUtils;
import org.apache.commons.lang3.StringUtils;
import java.io.IOException;
import java.net.URL;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.HashMap;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* YOLO模型Criteria 构建器
* @author dwj
* @date 2025/5/14
*/
public class Tensorflow2CriteriaBuilder implements CriteriaBuilderStrategy {
@Override
public Criteria<Image, DetectedObjects> buildCriteria(DetectorModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Map<String, Object> customParams = getDefaultConfig();
// 合并用户自定义参数(如有重复,覆盖默认默认值)
if (config.getCustomParams() != null) {
customParams.putAll(config.getCustomParams());
}
//解析synset
String synsetUrl = (String)customParams.get("synsetUrl");
String synsetPath = (String)customParams.get("synsetPath");
String synsetFileName = (String)customParams.get("synsetFileName");
Map<Integer, String> classes = null;
if(StringUtils.isNotBlank(synsetUrl)){
try {
classes = TensorflowSynsetUtils.loadSynset(new URL(synsetUrl));
} catch (IOException e) {
throw new DetectionException("加载synset异常", e);
}
}else if(StringUtils.isNotBlank(synsetPath)){
try {
if(!Files.exists(Paths.get(synsetPath))){
throw new DetectionException("synsetPath:" + synsetPath + "不存在");
}
classes = TensorflowSynsetUtils.loadSynset(Paths.get(synsetPath));
} catch (IOException e) {
throw new DetectionException("加载synset异常", e);
}
}else if(StringUtils.isNotBlank(synsetFileName)){
if(StringUtils.isBlank(config.getModelPath())){
throw new DetectionException("指定synsetFileName需同时指定modelPath");
}
try {
Path modelPath = Paths.get(config.getModelPath());
Path synset = modelPath.resolve(synsetFileName);
if(!Files.exists(synset)){
throw new DetectionException(synset.toAbsolutePath().toString() + " 不存在");
}
//模型同目录下存在synsetFileName
classes = TensorflowSynsetUtils.loadSynset(synset);
} catch (IOException e) {
throw new DetectionException("加载synset异常", e);
}
}else{
if(StringUtils.isBlank(config.getModelPath())){
throw new DetectionException("modelPath is null");
}
try {
Path modelPath = Paths.get(config.getModelPath());
synsetFileName = "mscoco_label_map.pbtxt";
Path synset = modelPath.resolve(synsetFileName);
if(!Files.exists(synset)){
throw new DetectionException(synset.toAbsolutePath().toString() + " 不存在");
}
//模型同目录下存在synsetFileName
classes = TensorflowSynsetUtils.loadSynset(synset);
} catch (IOException e) {
throw new DetectionException("加载synset异常", e);
}
}
customParams.put("classes", classes);
Criteria.Builder criteriaBuilder = Criteria.builder()
.optApplication(Application.CV.OBJECT_DETECTION)
.setTypes(Image.class, DetectedObjects.class)
.optModelPath(Paths.get(config.getModelPath()))
.optModelName("saved_model")
.optEngine("TensorFlow")
.optDevice(device)
.optTranslator(new TensorflowTranslator(customParams))
.optProgress(new ProgressBar());
if(config.getMaxBox() > 0){
criteriaBuilder.optArgument("maxBox", config.getMaxBox());
}
Criteria<Image, DetectedObjects> criteria = criteriaBuilder.build();
return criteria;
}
public Map<String, Object> getDefaultConfig(){
return new HashMap<>();
}
}

View File

@@ -0,0 +1,65 @@
package cn.smartjavaai.objectdetection.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.translator.YoloV8TranslatorFactory;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
import java.nio.file.Paths;
import java.util.HashMap;
import java.util.Map;
import java.util.Objects;
/**
* YOLO模型Criteria 构建器
* @author dwj
* @date 2025/5/14
*/
public class YoloCriteriaBuilder implements CriteriaBuilderStrategy {
@Override
public Criteria<Image, DetectedObjects> buildCriteria(DetectorModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Map<String, Object> customParams = getDefaultConfig();
// 合并用户自定义参数(如有重复,覆盖默认默认值)
if (config.getCustomParams() != null) {
customParams.putAll(config.getCustomParams());
}
Criteria.Builder criteriaBuilder = Criteria.builder()
.setTypes(Image.class, DetectedObjects.class)
//.optModelUrls("/Users/wenjie/Documents/develop/face_model/yolo")
.optModelPath(Paths.get(config.getModelPath()))
.optEngine("OnnxRuntime")
//.optOption("ortDevice", "TensorRT")
.optArguments(customParams)
.optDevice(device)
.optTranslatorFactory(new YoloV8TranslatorFactory())
.optProgress(new ProgressBar())
.optArgument("threshold", config.getThreshold() > 0 ? config.getThreshold() : DetectorConstant.DEFAULT_THRESHOLD);
if(config.getMaxBox() > 0){
criteriaBuilder.optArgument("maxBox", config.getMaxBox());
}
Criteria<Image, DetectedObjects> criteria = criteriaBuilder.build();
return criteria;
}
public Map<String, Object> getDefaultConfig(){
Map<String, Object> arguments = new HashMap<>();
// 添加默认参数
arguments.put("width", 640);
arguments.put("height", 640);
arguments.put("resize", true);
arguments.put("toTensor", true);
arguments.put("applyRatio", true);
return arguments;
}
}

View File

@@ -0,0 +1,72 @@
package cn.smartjavaai.objectdetection.enums;
/**
* 目标检测模型枚举
* @author dwj
* @date 2025/4/4
*/
public enum DetectorModelEnum {
// resnet50 系列
SSD_300_RESNET50("ai.djl.pytorch/ssd/0.0.1/ssd_300_resnet50"),
SSD_512_RESNET50_V1_VOC("ai.djl./ssd/0.0.1/ssd_512_resnet50_v1_voc"),
// vgg16 系列
SSD_512_VGG16_ATROUS_COCO("ai.djl.mxnet/ssd/0.0.1/ssd_512_vgg16_atrous_coco"),
SSD_300_VGG16_ATROUS_VOC("ai.djl.mxnet/ssd/0.0.1/ssd_300_vgg16_atrous_voc"),
// mobilenet 系列
SSD_512_MOBILENET1_VOC("ai.djl.mxnet/ssd/0.0.1/ssd_512_mobilenet1.0_voc"),
// YOLO 系列
YOLOV8N("ai.djl.pytorch/yolov8n/0.0.1/yolov8n"),
YOLO11N("ai.djl.pytorch/yolo11n/0.0.1/yolo11n"),
YOLOV5S("ai.djl.pytorch/yolo5s/0.0.1/yolov5s"),
YOLOV5S_ONNXRUNTIME("ai.djl.onnxruntime/yolo5s/0.0.1/yolo5s"),
YOLO("ai.djl.mxnet/yolo/0.0.1/yolo"),
// YOLOv3 变体
YOLO3_DARKNET_VOC_416("ai.djl.mxnet/yolo/0.0.1/yolo3_darknet_voc_416"),
YOLO3_MOBILENET_VOC_320("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_voc_320"),
YOLO3_MOBILENET_VOC_416("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_voc_41"),
YOLO3_DARKNET_COCO_320("ai.djl.mxnet/yolo/0.0.1/yolo3_darknet_coco_320"),
YOLO3_DARKNET_COCO_416("ai.djl.mxnet/yolo/0.0.1/yolo3_darknet_coco_416"),
YOLO3_DARKNET_COCO_608("ai.djl.mxnet/yolo/0.0.1/yolo3_darknet_coco_608"),
YOLO3_MOBILENET_COCO_320("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_coco_320"),
YOLO3_MOBILENET_COCO_416("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_coco_416"),
YOLO3_MOBILENET_COCO_608("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_coco_608"),
YOLOV12_OFFICIAL(""),
YOLOV8_OFFICIAL(""),
YOLOV8_CUSTOM(""),
YOLOV12_CUSTOM(""),
// TensorFlow 2.x 官方模型
TENSORFLOW2_OFFICIAL("");
/**
* 根据名称获取枚举 (忽略大小写和下划线变体)
*/
public static DetectorModelEnum fromName(String name) {
String formatted = name.trim().toUpperCase().replaceAll("[-_]", "");
for (DetectorModelEnum model : values()) {
if (model.name().replaceAll("_", "").equals(formatted)) {
return model;
}
}
throw new IllegalArgumentException("未知模型名称: " + name);
}
private final String modelUri;
DetectorModelEnum(String modelUri) {
this.modelUri = modelUri;
}
public String getModelUri() {
return modelUri;
}
}

View File

@@ -0,0 +1,51 @@
package cn.smartjavaai.objectdetection.enums;
/**
* 目标检测模型枚举
* @author dwj
* @date 2025/4/4
*/
public enum PersonDetectorModelEnum {
YOLOV8_PERSON("OnnxRuntime", 1280);
/**
* 模型输入尺寸
*/
private final int inputSize;
/**
* 模型引擎
*/
private final String engine;
PersonDetectorModelEnum(String engine, int inputSize) {
this.inputSize = inputSize;
this.engine = engine;
}
public int getInputSize() {
return inputSize;
}
public String getEngine() {
return engine;
}
/**
* 根据名称获取枚举 (忽略大小写和下划线变体)
*/
public static PersonDetectorModelEnum fromName(String name) {
String formatted = name.trim().toUpperCase().replaceAll("[-_]", "");
for (PersonDetectorModelEnum model : values()) {
if (model.name().replaceAll("_", "").equals(formatted)) {
return model;
}
}
throw new IllegalArgumentException("未知模型名称: " + name);
}
}

View File

@@ -0,0 +1,30 @@
package cn.smartjavaai.objectdetection.exception;
/**
* 目标检测异常
* @author dwj
* @date 2025/4/4
*/
public class DetectionException extends RuntimeException{
public DetectionException() {
super();
}
public DetectionException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) {
super(message, cause, enableSuppression, writableStackTrace);
}
public DetectionException(String message, Throwable cause) {
super(message, cause);
}
public DetectionException(String message) {
super(message);
}
public DetectionException(Throwable cause) {
super(cause);
}
}

View File

@@ -0,0 +1,270 @@
package cn.smartjavaai.objectdetection.model;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.FrameConverterUtil;
import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
import cn.smartjavaai.objectdetection.criteria.CriteriaBuilderFactory;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.vision.utils.CategoryMaskFilter;
import cn.smartjavaai.vision.utils.DetectedObjectsFilter;
import cn.smartjavaai.vision.utils.DetectorUtils;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.pool2.impl.GenericObjectPool;
import org.opencv.core.Mat;
import javax.imageio.ImageIO;
import java.awt.image.BufferedImage;
import java.io.*;
import java.nio.file.Paths;
import java.util.ArrayList;
import java.util.List;
import java.util.Objects;
/**
* 目标检测模型
* @author dwj
*/
@Slf4j
public class DetectorModel implements AutoCloseable{
private ZooModel<Image, DetectedObjects> model;
private GenericObjectPool<Predictor<Image, DetectedObjects>> predictorPool;
private DetectorModelConfig config;
private boolean fromFactory = false;
public void setFromFactory(boolean fromFactory) {
this.fromFactory = fromFactory;
}
public boolean isFromFactory() {
return fromFactory;
}
public void loadModel(DetectorModelConfig config){
if(Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型枚举");
}
Criteria<Image, DetectedObjects> criteria = CriteriaBuilderFactory.createCriteria(config);
this.config = config;
try {
model = criteria.loadModel();
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model));
int predictorPoolSize = config.getPredictorPoolSize();
if(config.getPredictorPoolSize() <= 0){
predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数
}
predictorPool.setMaxTotal(predictorPoolSize);
log.debug("当前设备: " + model.getNDManager().getDevice());
log.debug("当前引擎: " + Engine.getInstance().getEngineName());
log.debug("模型推理器线程池最大数量: " + predictorPoolSize);
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new DetectionException("模型加载失败", e);
}
}
/**
* 目标检测
* @param imagePath
* @return
* @throws Exception
*/
public DetectionResponse detect(String imagePath){
if(!FileUtils.isFileExists(imagePath)){
throw new DetectionException("图像文件不存在");
}
Image image = null;
try {
image = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
DetectedObjects detectedObjects = detect(image);
return DetectorUtils.convertToDetectionResponse(detectedObjects, image);
} catch (Exception e) {
throw new DetectionException(e);
} finally {
if (image != null){
((Mat)image.getWrappedImage()).release();
}
}
}
/**
* 目标检测-将检测结果绘制到原图
* @param imagePath
* @param outputPath
*/
public void detectAndDraw(String imagePath, String outputPath){
if(!FileUtils.isFileExists(imagePath)){
throw new DetectionException("图像文件不存在");
}
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
DetectedObjects detectedObjects = detect(img);
img.drawBoundingBoxes(detectedObjects);
ByteArrayOutputStream outputStream = new ByteArrayOutputStream();
// 调用 save 方法将 Image 写入字节流
img.save(new FileOutputStream(Paths.get(outputPath).toAbsolutePath().toString()), "png");
} catch (IOException e) {
throw new DetectionException(e);
} finally {
if (img != null){
((Mat)img.getWrappedImage()).release();
}
}
}
/**
* 目标检测
* @param imageData
* @return
*/
public DetectionResponse detect(byte[] imageData){
if(Objects.isNull(imageData)){
throw new DetectionException("图像无效");
}
try {
BufferedImage image = ImageIO.read(new ByteArrayInputStream(imageData));
return detect(image);
} catch (IOException e) {
throw new DetectionException("错误的图像", e);
}
}
/**
* 目标检测
* @param image
* @return
*/
public DetectionResponse detect(BufferedImage image){
if(!ImageUtils.isImageValid(image)){
throw new DetectionException("图像无效");
}
Image img = null;
try {
img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
DetectedObjects detectedObjects = detect(img);
return DetectorUtils.convertToDetectionResponse(detectedObjects, img);
} catch (Exception e) {
throw new DetectionException(e);
} finally {
if (img != null) {
((Mat)img.getWrappedImage()).release();
}
}
}
/**
* 目标检测-将检测结果绘制到原图
* @param sourceImage
* @return
*/
public BufferedImage detectAndDraw(BufferedImage sourceImage){
if(!ImageUtils.isImageValid(sourceImage)){
throw new DetectionException("图像无效");
}
Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage));
DetectedObjects detectedObjects = detect(img);
img.drawBoundingBoxes(detectedObjects);
try {
ByteArrayOutputStream outputStream = new ByteArrayOutputStream();
// 调用 save 方法将 Image 写入字节流
img.save(outputStream, "png");
// 将字节流转换为 BufferedImage
byte[] imageBytes = outputStream.toByteArray();
return ImageIO.read(new ByteArrayInputStream(imageBytes));
} catch (IOException e) {
throw new DetectionException("导出图片失败", e);
} finally {
if (img != null) {
((Mat)img.getWrappedImage()).release();
}
}
}
/**
* 目标检测
* @param image
* @return
*/
public DetectedObjects detect(Image image){
Predictor<Image, DetectedObjects> predictor = null;
try {
predictor = predictorPool.borrowObject();
DetectedObjects detectedObjects = predictor.predict(image);
if(CollectionUtils.isNotEmpty(config.getAllowedClasses())
&& Objects.nonNull(detectedObjects) && detectedObjects.getNumberOfObjects() > 0){
DetectedObjectsFilter detectedObjectsFilter = new DetectedObjectsFilter(config.getAllowedClasses(), config.getThreshold(),config.getTopK());
detectedObjects = detectedObjectsFilter.filter(detectedObjects);
}
return detectedObjects;
} catch (Exception e) {
throw new DetectionException("目标检测错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
log.debug("释放资源");
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
public GenericObjectPool<Predictor<Image, DetectedObjects>> getPool() {
return predictorPool;
}
/**
* 显式释放资源
*/
@Override
public void close() {
if (fromFactory) {
ObjectDetectionModelFactory.removeFromCache(config.getModelEnum());
}
try {
if (predictorPool != null) {
predictorPool.close();
}
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (model != null) {
model.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,87 @@
package cn.smartjavaai.objectdetection.model;
import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import lombok.extern.slf4j.Slf4j;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* 目标检测 模型工厂
* @author dwj
*/
@Slf4j
public class ObjectDetectionModelFactory {
// 使用 volatile 和双重检查锁定来确保线程安全的单例模式
private static volatile ObjectDetectionModelFactory instance;
private static final ConcurrentHashMap<DetectorModelEnum, DetectorModel> modelMap = new ConcurrentHashMap<>();
static{
log.debug("缓存目录:{}", Config.getCachePath());
}
// 私有构造函数,防止外部创建实例
private ObjectDetectionModelFactory() {}
// 双重检查锁定的单例方法
public static ObjectDetectionModelFactory getInstance() {
if (instance == null) {
synchronized (ObjectDetectionModelFactory.class) {
if (instance == null) {
instance = new ObjectDetectionModelFactory();
}
}
}
return instance;
}
/**
* 获取模型(通过配置)
* @param config
* @return
*/
public DetectorModel getModel(DetectorModelConfig config) {
if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型");
}
return modelMap.computeIfAbsent(config.getModelEnum(), k -> {
DetectorModel model = new DetectorModel();
model.loadModel(config);
model.setFromFactory(true);
return model;
});
}
/**
* 获取默认模型
* @return
*/
public DetectorModel getModel() {
// 初始化默认配置
DetectorModelConfig config = new DetectorModelConfig();
config.setModelEnum(DetectorModelEnum.YOLO11N);
return getModel(config);
}
/**
* 关闭所有已加载的模型
*/
public void closeAll() {
modelMap.values().forEach(DetectorModel::close);
modelMap.clear();
}
/**
* 移除缓存的模型
* @param modelEnum
*/
public static void removeFromCache(DetectorModelEnum modelEnum) {
modelMap.remove(modelEnum);
}
}

View File

@@ -0,0 +1,160 @@
package cn.smartjavaai.objectdetection.model.person;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import cn.smartjavaai.common.cv.SmartImageFactory;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.obb.config.ObbDetModelConfig;
import cn.smartjavaai.obb.criteria.ObbDetCriteriaFactory;
import cn.smartjavaai.obb.entity.ObbResult;
import cn.smartjavaai.obb.exception.ObbDetException;
import cn.smartjavaai.objectdetection.config.PersonDetModelConfig;
import cn.smartjavaai.objectdetection.criteria.PersonDetCriteriaFactory;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.vision.utils.DetectedObjectsFilter;
import cn.smartjavaai.vision.utils.DetectorUtils;
import cn.smartjavaai.vision.utils.ObbResultFilter;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.pool2.impl.GenericObjectPool;
import org.opencv.core.Mat;
import java.io.ByteArrayOutputStream;
import java.io.FileOutputStream;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Paths;
import java.util.Objects;
/**
* 行人检测模型
* @author dwj
*/
@Slf4j
public class CommonPersonDetModel implements PersonDetModel {
private PersonDetModelConfig config;
private ZooModel<Image, DetectedObjects> model;
private GenericObjectPool<Predictor<Image, DetectedObjects>> predictorPool;
@Override
public void loadModel(PersonDetModelConfig config) {
if(Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型枚举");
}
Criteria<Image, DetectedObjects> criteria = PersonDetCriteriaFactory.createCriteria(config);
this.config = config;
try {
model = criteria.loadModel();
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model));
int predictorPoolSize = config.getPredictorPoolSize();
if(config.getPredictorPoolSize() <= 0){
predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数
}
predictorPool.setMaxTotal(predictorPoolSize);
log.debug("当前设备: " + model.getNDManager().getDevice());
log.debug("当前引擎: " + Engine.getInstance().getEngineName());
log.debug("模型推理器线程池最大数量: " + predictorPoolSize);
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new DetectionException("模型加载失败", e);
}
}
@Override
public R<DetectionResponse> detect(Image image) {
DetectedObjects detectedObjects = detectCore(image);
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image);
return R.ok(detectionResponse);
}
/**
* 模型核心推理方法
* @param image
* @return
*/
@Override
public DetectedObjects detectCore(Image image) {
Predictor<Image, DetectedObjects> predictor = null;
try {
predictor = predictorPool.borrowObject();
DetectedObjects detectedObjects = predictor.predict(image);
//过滤
if(Objects.nonNull(detectedObjects) && detectedObjects.getNumberOfObjects() > 0){
DetectedObjectsFilter detectedObjectsFilter = new DetectedObjectsFilter(config.getAllowedClasses(), config.getTopK());
detectedObjects = detectedObjectsFilter.filter(detectedObjects);
}
return detectedObjects;
} catch (Exception e) {
throw new DetectionException("行人检测错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
log.debug("释放资源");
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public R<DetectionResponse> detectAndDraw(Image image) {
DetectedObjects detectedObjects = detectCore(image);
image.drawBoundingBoxes(detectedObjects);
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image);
detectionResponse.setDrawnImage(image);
return R.ok(detectionResponse);
}
@Override
public R<DetectionResponse> detectAndDraw(String imagePath, String outputPath) {
try {
Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath));
DetectedObjects detectedObjects = detectCore(img);
img.drawBoundingBoxes(detectedObjects);
ByteArrayOutputStream outputStream = new ByteArrayOutputStream();
// 调用 save 方法将 Image 写入字节流
img.save(new FileOutputStream(Paths.get(outputPath).toAbsolutePath().toString()), "png");
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, img);
return R.ok(detectionResponse);
} catch (IOException e) {
throw new ObbDetException(e);
}
}
@Override
public void close() throws Exception {
try {
if (predictorPool != null) {
predictorPool.close();
}
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (model != null) {
model.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,62 @@
package cn.smartjavaai.objectdetection.model.person;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.obb.config.ObbDetModelConfig;
import cn.smartjavaai.obb.entity.ObbResult;
import cn.smartjavaai.objectdetection.config.PersonDetModelConfig;
/**
* 行人检测模型
* @author dwj
*/
public interface PersonDetModel extends AutoCloseable{
/**
* 加载模型
* @param config
*/
void loadModel(PersonDetModelConfig config);
/**
* 旋转框
* @param image
* @return
*/
default R<DetectionResponse> detect(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 行人检测 核心方法
* @param image
* @return
*/
default DetectedObjects detectCore(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 行人检测并绘制
* @param image
* @return
*/
default R<DetectionResponse> detectAndDraw(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 行人检测并绘制
* @param imagePath
* @param outputPath
* @return
*/
default R<DetectionResponse> detectAndDraw(String imagePath, String outputPath){
throw new UnsupportedOperationException("默认不支持该功能");
}
}

View File

@@ -0,0 +1,16 @@
package cn.smartjavaai.objectdetection.stream;
import ai.djl.modality.cv.Image;
import cn.smartjavaai.common.entity.DetectionInfo;
import java.util.List;
/**
* 视频流目标检测监听器
* @author dwj
*/
public interface StreamDetectionListener {
void onObjectDetected(List<DetectionInfo> detectionInfoList, Image image);
}

View File

@@ -0,0 +1,280 @@
package cn.smartjavaai.objectdetection.stream;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.output.Mask;
import ai.djl.modality.cv.output.Rectangle;
import cn.hutool.core.lang.UUID;
import cn.smartjavaai.common.entity.*;
import cn.smartjavaai.common.enums.VideoSourceType;
import cn.smartjavaai.common.utils.FrameConverterUtil;
import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.objectdetection.model.DetectorModel;
import cn.smartjavaai.vision.utils.DetectorUtils;
import lombok.extern.slf4j.Slf4j;
import nu.pattern.OpenCV;
import org.apache.commons.lang3.StringUtils;
import org.bytedeco.javacv.*;
import org.bytedeco.opencv.global.opencv_imgcodecs;
import org.opencv.core.Core;
import org.opencv.core.Mat;
import org.opencv.imgcodecs.Imgcodecs;
import java.util.*;
import java.util.concurrent.*;
/**
* 视频流目标检测器
* @author dwj
*/
@Slf4j
public class StreamDetector implements AutoCloseable{
static {
OpenCV.loadLocally();
}
private DetectorModel detectorModel;
private String streamUrl;
private ExecutorService grabberExecutor; // 专门抓帧的线程
private ExecutorService processorExecutor; // 专门处理帧的线程
private int frameDetectionInterval = 1;
private int repeatGap = 5; // 秒
private volatile boolean isRunning;
private FrameGrabber grabber;
private StreamDetectionListener listener;
private OpenCVFrameConverter.ToOrgOpenCvCoreMat converterToMat;
private VideoSourceType sourceType = VideoSourceType.STREAM; // 默认流
private int cameraIndex = 0; // 默认第一个摄像头
private Map<String, Long> lastDetectTime = new ConcurrentHashMap<>();
private BlockingQueue<Frame> frameQueue = new LinkedBlockingQueue<>(100);
public static Builder builder() { return new Builder(); }
private StreamDetector(Builder builder) {
this.detectorModel = builder.detectorModel;
this.streamUrl = builder.streamUrl;
this.frameDetectionInterval = builder.frameDetectionInterval;
this.listener = builder.listener;
this.sourceType = builder.sourceType;
this.cameraIndex = builder.cameraIndex;
this.converterToMat = new OpenCVFrameConverter.ToOrgOpenCvCoreMat();
}
private void initializeGrabber() throws FrameGrabber.Exception {
if (sourceType == VideoSourceType.CAMERA) {
// 使用用户传入的摄像头索引
grabber = new OpenCVFrameGrabber(cameraIndex);
} else {
grabber = new FFmpegFrameGrabber(streamUrl);
if (sourceType == VideoSourceType.STREAM) {
grabber.setOption("rtsp_transport", "tcp");
grabber.setOption("buffer_size", "1024000");
grabber.setOption("stimeout", "20000000");
grabber.setOption("max_delay", "500000");
}
}
grabber.start();
}
public void startDetection() {
if (isRunning) return;
isRunning = true;
// 初始化抓帧线程池
if (grabberExecutor == null) grabberExecutor = Executors.newSingleThreadExecutor();
// 初始化帧处理线程池
if (processorExecutor == null) processorExecutor = Executors.newSingleThreadExecutor();
// 初始化抓帧线程
grabberExecutor.submit(() -> {
try {
initializeGrabber();
processFrames();
} catch (Exception e) {
log.error("视频流处理异常", e);
} finally {
release();
}
});
log.info("视频流处理已启动");
// 初始化队列处理线程:解决回调比较耗时,导致线程池爆满
startFrameProcessor();
}
/**
* 负责抓取视频帧到队列
*/
private void processFrames() {
int frameCount = 0;
while (isRunning) {
try {
Frame frame = grabber.grab();
if (frame == null || frame.image == null) continue;
frameCount++;
if (frameCount % frameDetectionInterval != 0) continue;
Frame currentFrame = frame.clone();
frameQueue.offer(currentFrame); // 队列满则丢弃,可改为 put 阻塞
} catch (Exception e) {
log.error("抓取视频帧异常", e);
if (e instanceof FFmpegFrameGrabber.Exception) reconnect();
}
}
}
/**
* 负责处理视频帧
*/
private void startFrameProcessor() {
processorExecutor.submit(() -> {
log.info("帧处理线程已启动");
while (isRunning || !frameQueue.isEmpty()) {
try {
Frame frame = frameQueue.poll(100, TimeUnit.MILLISECONDS);
if (frame != null) processFrame(frame);
} catch (InterruptedException e) {
e.printStackTrace();
} catch (Exception e) {
log.error("帧处理异常", e);
}
}
});
}
private void processFrame(Frame frame) {
Mat mat = null;
try {
mat = converterToMat.convert(frame);
if (mat == null) return;
Image image = ImageFactory.getInstance().fromImage(mat);
DetectedObjects detectedObjects = detectorModel.detect(image);
// log.debug("检测结果:{}", detectedObjects.toString());
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image);
List<DetectionInfo> filtered = filterRepeatedObjects(detectionResponse);
if (!filtered.isEmpty() && listener != null) {
listener.onObjectDetected(filtered, image); // 同帧多物体一次回调
}
} catch (Throwable e) {
e.printStackTrace();
log.error("单帧处理异常", e);
} finally {
if (mat != null) mat.release();
}
}
private List<DetectionInfo> filterRepeatedObjects(DetectionResponse response) {
List<DetectionInfo> result = new ArrayList<>();
long now = System.currentTimeMillis();
for (DetectionInfo info : response.getDetectionInfoList()) {
String name = info.getObjectDetInfo().getClassName();
Long last = lastDetectTime.get(name);
if (last == null || (now - last) > repeatGap * 1000) {
lastDetectTime.put(name, now);
result.add(info);
}
}
return result;
}
private void reconnect() {
log.info("尝试重新连接视频流");
try {
release();
Thread.sleep(5000);
initializeGrabber();
} catch (Exception e) {
log.error("重新连接RTSP流失败", e);
}
}
public void stopDetection() { isRunning = false; }
private void release() {
if (grabber != null) {
try { grabber.stop(); grabber.release(); }
catch (FrameGrabber.Exception e) { log.error("释放Grabber失败", e); }
}
}
@Override
public void close() {
stopDetection();
if (grabberExecutor != null) grabberExecutor.shutdownNow();
if (processorExecutor != null) processorExecutor.shutdownNow();
release();
}
public static class Builder {
private DetectorModel detectorModel;
private String streamUrl;
private ExecutorService executorService;
private int frameDetectionInterval = 1;
private StreamDetectionListener listener;
private VideoSourceType sourceType = VideoSourceType.STREAM; // 默认流
private int cameraIndex = 0; // 默认第一个摄像头
public Builder detectorModel(DetectorModel m) { this.detectorModel = m; return this; }
public Builder streamUrl(String url) { this.streamUrl = url; return this; }
public Builder executorService(ExecutorService es) { this.executorService = es; return this; }
public Builder listener(StreamDetectionListener listener) { this.listener = listener; return this; }
public Builder sourceType(VideoSourceType sourceType) {
this.sourceType = sourceType;
return this;
}
public Builder cameraIndex(int cameraIndex) {
this.cameraIndex = cameraIndex;
return this;
}
public Builder frameDetectionInterval(int interval) {
if (interval < 1) throw new IllegalArgumentException("frameDetectionInterval >= 1");
this.frameDetectionInterval = interval;
return this;
}
public StreamDetector build() {
if (detectorModel == null) {
throw new DetectionException("detectorModel 不能为空");
}
if (sourceType == null) {
throw new DetectionException("sourceType 不能为空");
}
// 根据 sourceType 校验不同的参数
switch (sourceType) {
case STREAM:
case FILE:
if (StringUtils.isBlank(streamUrl)) {
throw new DetectionException("streamUrl 不能为空");
}
break;
case CAMERA:
if (cameraIndex < 0) {
throw new DetectionException("cameraIndex 必须 >= 0");
}
break;
default:
throw new DetectionException("不支持的视频源类型: " + sourceType);
}
if (executorService == null) {
executorService = Executors.newFixedThreadPool(2); // 至少2个线程
}
return new StreamDetector(this);
}
}
}

View File

@@ -0,0 +1,111 @@
package cn.smartjavaai.objectdetection.translator;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.output.Rectangle;
import ai.djl.modality.cv.util.NDImageUtils;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.types.DataType;
import ai.djl.translate.NoBatchifyTranslator;
import ai.djl.translate.TranslatorContext;
import java.io.IOException;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
import java.util.Objects;
/**
* @author dwj
* @date 2025/8/18
*/
public class TensorflowTranslator implements NoBatchifyTranslator<Image, DetectedObjects> {
private Map<Integer, String> classes;
private int maxBoxes;
private float threshold;
public TensorflowTranslator(Map<String, ?> arguments) {
maxBoxes =
arguments.containsKey("maxBoxes")
? Integer.parseInt(arguments.get("maxBoxes").toString())
: 10;
threshold =
arguments.containsKey("threshold")
? Float.parseFloat(arguments.get("threshold").toString())
: 0.7f;
classes = (Map<Integer, String>)arguments.get("classes");
}
/** {@inheritDoc} */
@Override
public NDList processInput(TranslatorContext ctx, Image input) {
// input to tf object-detection models is a list of tensors, hence NDList
NDArray array = input.toNDArray(ctx.getNDManager(), Image.Flag.COLOR);
// optionally resize the image for faster processing
array = NDImageUtils.resize(array, 224);
// tf object-detection models expect 8 bit unsigned integer tensor
array = array.toType(DataType.UINT8, true);
array = array.expandDims(0); // tf object-detection models expect a 4 dimensional input
return new NDList(array);
}
/** {@inheritDoc} */
@Override
public void prepare(TranslatorContext ctx) throws IOException {
}
/** {@inheritDoc} */
@Override
public DetectedObjects processOutput(TranslatorContext ctx, NDList list) {
// output of tf object-detection models is a list of tensors, hence NDList in djl
// output NDArray order in the list are not guaranteed
int[] classIds = null;
float[] probabilities = null;
NDArray boundingBoxes = null;
for (NDArray array : list) {
if ("detection_boxes".equals(array.getName())) {
boundingBoxes = array.get(0);
} else if ("detection_scores".equals(array.getName())) {
probabilities = array.get(0).toFloatArray();
} else if ("detection_classes".equals(array.getName())) {
// class id is between 1 - number of classes
classIds = array.get(0).toType(DataType.INT32, true).toIntArray();
}
}
Objects.requireNonNull(classIds);
Objects.requireNonNull(probabilities);
Objects.requireNonNull(boundingBoxes);
List<String> retNames = new ArrayList<>();
List<Double> retProbs = new ArrayList<>();
List<BoundingBox> retBB = new ArrayList<>();
// result are already sorted
for (int i = 0; i < Math.min(classIds.length, maxBoxes); ++i) {
int classId = classIds[i];
double probability = probabilities[i];
// classId starts from 1, -1 means background
if (classId > 0 && probability > threshold) {
String className = classes.getOrDefault(classId, "#" + classId);
float[] box = boundingBoxes.get(i).toFloatArray();
float yMin = box[0];
float xMin = box[1];
float yMax = box[2];
float xMax = box[3];
Rectangle rect = new Rectangle(xMin, yMin, xMax - xMin, yMax - yMin);
retNames.add(className);
retProbs.add(probability);
retBB.add(rect);
}
}
return new DetectedObjects(retNames, retProbs, retBB);
}
}

View File

@@ -0,0 +1,457 @@
package cn.smartjavaai.objectdetection.translator;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.output.Rectangle;
import ai.djl.modality.cv.transform.*;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import ai.djl.translate.*;
import cn.smartjavaai.common.utils.LetterBoxUtils;
import java.nio.file.Files;
import java.nio.file.Paths;
import java.util.*;
/**
* YoloV8 行人检测 Translator
*/
public class YoloV8PersonDetTranslator implements Translator<Image, DetectedObjects> {
private int maxBoxes;
private YoloOutputType yoloOutputLayerType;
private float nmsThreshold;
protected float threshold;
// private BaseImageTranslator.SynsetLoader synsetLoader;
protected List<String> classes;
protected boolean applyRatio;
protected boolean removePadding;
protected Pipeline pipeline;
private Image.Flag flag;
private Batchifier batchifier;
protected int width;
protected int height;
/**
* Constructs an ImageTranslator with the provided builder.
*
* @param builder the data to build with
*/
protected YoloV8PersonDetTranslator(Builder builder) {
this.yoloOutputLayerType = builder.outputType;
this.nmsThreshold = builder.nmsThreshold;
maxBoxes = builder.maxBox;
this.threshold = builder.threshold;
// this.synsetLoader = builder.synsetLoader;
this.applyRatio = builder.applyRatio;
this.removePadding = builder.removePadding;
this.flag = builder.flag;
this.pipeline = builder.pipeline;
this.batchifier = builder.batchifier;
this.width = builder.width;
this.height = builder.height;
classes = Arrays.asList("person");
}
/**
* Creates a builder to build a {@code YoloV8Translator} with specified arguments.
*
* @return a new builder
*/
public static Builder builder() {
return new Builder();
}
/**
* Creates a builder to build a {@code YoloV8Translator} with specified arguments.
*
* @param arguments arguments to specify builder options
* @return a new builder
*/
public static Builder builder(Map<String, ?> arguments) {
Builder builder = new Builder();
builder.configPreProcess(arguments);
builder.configPostProcess(arguments);
return builder;
}
@Override
public NDList processInput(TranslatorContext ctx, Image input) throws Exception {
NDManager manager = ctx.getNDManager();
NDArray array = input.toNDArray(manager, Image.Flag.COLOR);
//Letter box resize 640x640 with padding (保持比例,补边缘)
LetterBoxUtils.ResizeResult letterBoxResult = LetterBoxUtils.letterbox(manager, array, width, height, 114f, LetterBoxUtils.PaddingPosition.CENTER);
array = letterBoxResult.image;
// 转为 float32 且归一化到 0~1
array = array.toType(DataType.FLOAT32, false).div(255f); // HWC
// HWC -> CHW
array = array.transpose(2, 0, 1); // CHW
ctx.setAttachment("width", input.getWidth());
ctx.setAttachment("height", input.getHeight());
ctx.setAttachment("processedWidth", width);
ctx.setAttachment("processedHeight", height);
ctx.setAttachment("scale", letterBoxResult.r);
return new NDList(array);
}
@Override
public DetectedObjects processOutput(TranslatorContext ctx, NDList list) throws Exception {
//原图宽高
int imageWidth = (Integer) ctx.getAttachment("width");
int imageHeight = (Integer) ctx.getAttachment("height");
float scale = (Float) ctx.getAttachment("scale");
switch (yoloOutputLayerType) {
case DETECT:
return processFromDetectOutput();
case AUTO:
if (list.get(0).getShape().dimension() > 2) {
return processFromDetectOutput();
} else {
return processFromBoxOutput(imageWidth, imageHeight, list, scale);
}
case BOX:
default:
return processFromBoxOutput(imageWidth, imageHeight, list, scale);
}
}
/** {@inheritDoc} */
protected DetectedObjects processFromBoxOutput(int origImageWidth, int origImageHeight, NDList list, float scale) {
NDArray rawResult = list.get(0);
NDArray reshapedResult = rawResult.transpose();
Shape shape = reshapedResult.getShape();
float[] buf = reshapedResult.toFloatArray();
int numberRows = Math.toIntExact(shape.get(0));
int nClasses = Math.toIntExact(shape.get(1));
int padding = nClasses - classes.size();
System.out.println(Arrays.toString(reshapedResult.get(0).toFloatArray()));
if (padding != 0 && padding != 4) {
throw new IllegalStateException(
"Expected classes: " + (nClasses - 4) + ", got " + classes.size());
}
ArrayList<Rectangle> boxes = new ArrayList<>();
ArrayList<Float> scores = new ArrayList<>();
ArrayList<Integer> classIds = new ArrayList<>();
// reverse order search in heap; searches through #maxBoxes for optimization when set
for (int i = numberRows - 1; i > numberRows - maxBoxes; --i) {
int index = i * nClasses;
float maxClassProb = buf[index + 4];
if (maxClassProb > threshold) {
float xPos = buf[index]; // center x
float yPos = buf[index + 1]; // center y
float w = buf[index + 2];
float h = buf[index + 3];
Rectangle rect =
new Rectangle(Math.max(0, xPos - w / 2), Math.max(0, yPos - h / 2), w, h);
scores.add(maxClassProb);
classIds.add(0);
boxes.add(rect);
}
}
return nms(origImageWidth, origImageHeight, boxes, classIds, scores, scale);
}
private DetectedObjects processFromDetectOutput() {
throw new UnsupportedOperationException(
"detect layer output is not supported yet, check correct YoloV5 export format");
}
protected DetectedObjects nms(
int origImageWidth,
int origImageHeight,
List<Rectangle> boxes,
List<Integer> classIds,
List<Float> scores, float scale) {
List<String> retClasses = new ArrayList<>();
List<Double> retProbs = new ArrayList<>();
List<BoundingBox> retBB = new ArrayList<>();
for (int classId = 0; classId < classes.size(); classId++) {
List<Rectangle> r = new ArrayList<>();
List<Double> s = new ArrayList<>();
List<Integer> map = new ArrayList<>();
for (int j = 0; j < classIds.size(); ++j) {
if (classIds.get(j) == classId) {
r.add(boxes.get(j));
s.add(scores.get(j).doubleValue());
map.add(j);
}
}
if (r.isEmpty()) {
continue;
}
List<Integer> nms = Rectangle.nms(r, s, nmsThreshold);
for (int index : nms) {
int pos = map.get(index);
int id = classIds.get(pos);
retClasses.add(classes.get(id));
retProbs.add(scores.get(pos).doubleValue());
// Rectangle rect = boxes.get(pos);
Rectangle rect = boxes.get(pos);
if (removePadding) {
rect =
LetterBoxUtils.restoreBox(rect, scale, origImageWidth, origImageHeight, width, height);
} else if (applyRatio) {
rect =
new Rectangle(
rect.getX() / width,
rect.getY() / height,
rect.getWidth() / width,
rect.getHeight() / height);
}
retBB.add(rect);
}
}
return new DetectedObjects(retClasses, retProbs, retBB);
}
public static class Builder {
private int maxBox = 8400;
YoloOutputType outputType;
float nmsThreshold;
protected float threshold = 0.2F;
protected boolean applyRatio;
protected boolean removePadding;
protected int width = 224;
protected int height = 224;
protected Image.Flag flag;
protected Pipeline pipeline;
protected Batchifier batchifier;
public Builder() {
this.outputType = YoloOutputType.AUTO;
this.nmsThreshold = 0.45F;
}
public Builder optOutputType(YoloOutputType outputType) {
this.outputType = outputType;
return this;
}
public Builder optNmsThreshold(float nmsThreshold) {
this.nmsThreshold = nmsThreshold;
return this;
}
/**
* Builds the translator.
*
* @return the new translator
*/
public YoloV8PersonDetTranslator build() {
if (pipeline == null) {
addTransform(
array -> array.transpose(2, 0, 1).toType(DataType.FLOAT32, false).div(255));
}
// validate();
return new YoloV8PersonDetTranslator(this);
}
protected Builder self() {
return this;
}
public Builder addTransform(Transform transform) {
if (this.pipeline == null) {
this.pipeline = new Pipeline();
}
this.pipeline.add(transform);
return this.self();
}
public Builder optApplyRatio(boolean value) {
this.applyRatio = value;
return this.self();
}
public Builder optFlag(Image.Flag flag) {
this.flag = flag;
return this.self();
}
public Builder setPipeline(Pipeline pipeline) {
this.pipeline = pipeline;
return this.self();
}
public Builder setImageSize(int width, int height) {
this.width = width;
this.height = height;
return this.self();
}
public Builder optBatchifier(Batchifier batchifier) {
this.batchifier = batchifier;
return this.self();
}
public Builder optThreshold(float threshold) {
this.threshold = threshold;
return this.self();
}
/** {@inheritDoc} */
protected void configPostProcess(Map<String, ?> arguments) {
if (ArgumentsUtil.booleanValue(arguments, "optApplyRatio") || ArgumentsUtil.booleanValue(arguments, "applyRatio")) {
this.optApplyRatio(true);
}
this.threshold = ArgumentsUtil.floatValue(arguments, "threshold", 0.2F);
String centerFit = ArgumentsUtil.stringValue(arguments, "centerFit", "false");
this.removePadding = "true".equals(centerFit);
String type = ArgumentsUtil.stringValue(arguments, "outputType", "AUTO");
this.outputType = YoloOutputType.valueOf(type.toUpperCase(Locale.ENGLISH));
this.nmsThreshold = ArgumentsUtil.floatValue(arguments, "nmsThreshold", 0.45F);
maxBox = ArgumentsUtil.intValue(arguments, "maxBox", 8400);
}
protected void configPreProcess(Map<String, ?> arguments) {
if (this.pipeline == null) {
this.pipeline = new Pipeline();
}
this.width = ArgumentsUtil.intValue(arguments, "width", 224);
this.height = ArgumentsUtil.intValue(arguments, "height", 224);
if (arguments.containsKey("flag")) {
this.flag = Image.Flag.valueOf(arguments.get("flag").toString());
}
String pad = ArgumentsUtil.stringValue(arguments, "pad", "false");
if ("true".equals(pad)) {
this.addTransform(new Pad(0.0));
} else if (!"false".equals(pad)) {
double padding = Double.parseDouble(pad);
this.addTransform(new Pad(padding));
}
String resize = ArgumentsUtil.stringValue(arguments, "resize", "false");
int w;
int shortEdge;
if ("true".equals(resize)) {
this.addTransform(new Resize(this.width, this.height));
} else if (!"false".equals(resize)) {
String[] tokens = resize.split("\\s*,\\s*");
w = (int)Double.parseDouble(tokens[0]);
if (tokens.length > 1) {
shortEdge = (int)Double.parseDouble(tokens[1]);
} else {
shortEdge = w;
}
Image.Interpolation interpolation;
if (tokens.length > 2) {
interpolation = Image.Interpolation.valueOf(tokens[2]);
} else {
interpolation = Image.Interpolation.BILINEAR;
}
this.addTransform(new Resize(w, shortEdge, interpolation));
}
String resizeShort = ArgumentsUtil.stringValue(arguments, "resizeShort", "false");
if ("true".equals(resizeShort)) {
w = Math.max(this.width, this.height);
this.addTransform(new ResizeShort(w));
} else if (!"false".equals(resizeShort)) {
String[] tokens = resizeShort.split("\\s*,\\s*");
shortEdge = (int)Double.parseDouble(tokens[0]);
int longEdge;
if (tokens.length > 1) {
longEdge = (int)Double.parseDouble(tokens[1]);
} else {
longEdge = -1;
}
Image.Interpolation interpolation;
if (tokens.length > 2) {
interpolation = Image.Interpolation.valueOf(tokens[2]);
} else {
interpolation = Image.Interpolation.BILINEAR;
}
this.addTransform(new ResizeShort(shortEdge, longEdge, interpolation));
}
if (ArgumentsUtil.booleanValue(arguments, "centerCrop", false)) {
this.addTransform(new CenterCrop(this.width, this.height));
}
if (ArgumentsUtil.booleanValue(arguments, "centerFit")) {
this.addTransform(new CenterFit(this.width, this.height));
}
if (ArgumentsUtil.booleanValue(arguments, "toTensor", true)) {
this.addTransform(new ToTensor());
}
String normalize = ArgumentsUtil.stringValue(arguments, "normalize", "false");
if ("true".equals(normalize)) {
float[] MEAN = new float[]{0.485F, 0.456F, 0.406F};
float[] STD = new float[]{0.229F, 0.224F, 0.225F};
this.addTransform(new Normalize(MEAN, STD));
} else if (!"false".equals(normalize)) {
String[] tokens = normalize.split("\\s*,\\s*");
if (tokens.length != 6) {
throw new IllegalArgumentException("Invalid normalize value: " + normalize);
}
float[] mean = new float[]{Float.parseFloat(tokens[0]), Float.parseFloat(tokens[1]), Float.parseFloat(tokens[2])};
float[] std = new float[]{Float.parseFloat(tokens[3]), Float.parseFloat(tokens[4]), Float.parseFloat(tokens[5])};
this.addTransform(new Normalize(mean, std));
}
String range = (String)arguments.get("range");
if ("0,1".equals(range)) {
this.addTransform((a) -> {
return a.div(255.0F);
});
} else if ("-1,1".equals(range)) {
this.addTransform((a) -> {
return a.div(128.0F).sub(1);
});
}
if (arguments.containsKey("batchifier")) {
this.batchifier = Batchifier.fromString((String)arguments.get("batchifier"));
}
}
}
public static enum YoloOutputType {
BOX,
DETECT,
AUTO;
private YoloOutputType() {
}
}
}

View File

@@ -0,0 +1,49 @@
package cn.smartjavaai.semseg.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.semseg.enums.SemSegModelEnum;
import lombok.Data;
import java.util.List;
/**
* 语义分割模型参数配置
*
* @author dwj
*/
@Data
public class SemSegModelConfig extends ModelConfig {
/**
* 模型
*/
private SemSegModelEnum modelEnum;
/**
* 模型路径
*/
private String modelPath;
/**
* 允许的分类列表
*/
private List<String> allowedClasses;
public SemSegModelConfig() {
}
public SemSegModelConfig(SemSegModelEnum modelEnum, DeviceEnum device) {
this.modelEnum = modelEnum;
setDevice(device);
}
public SemSegModelConfig(SemSegModelEnum modelEnum) {
this.modelEnum = modelEnum;
}
}

View File

@@ -0,0 +1,48 @@
package cn.smartjavaai.semseg.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.CategoryMask;
import ai.djl.modality.cv.translator.SemanticSegmentationTranslatorFactory;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.semseg.config.SemSegModelConfig;
import cn.smartjavaai.semseg.enums.SemSegModelEnum;
import org.apache.commons.lang3.StringUtils;
import java.nio.file.Paths;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* 语义分割Criteria工厂
* @author dwj
*/
public class SemSegCriteriaFactory {
public static Criteria<Image, CategoryMask> createCriteria(SemSegModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Criteria<Image, CategoryMask> criteria = null;
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
params.putAll(config.getCustomParams());
if(config.getModelEnum() == SemSegModelEnum.DEEPLABV3){
criteria =
Criteria.builder()
.setTypes(Image.class, CategoryMask.class)
.optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null :
config.getModelEnum().getModelUri())
.optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null)
.optTranslatorFactory(new SemanticSegmentationTranslatorFactory())
.optEngine("PyTorch")
.optDevice(device)
.optProgress(new ProgressBar())
.build();
}
return criteria;
}
}

View File

@@ -0,0 +1,17 @@
package cn.smartjavaai.semseg.entity;
import lombok.Data;
/**
* 检测参数
* @author dwj
*/
@Data
public class DetectParams {
/**
* 置信度阈值
*/
private float threshold = 0.3f;
}

View File

@@ -0,0 +1,34 @@
package cn.smartjavaai.semseg.enums;
/**
* 语义分割模型枚举
* @author dwj
*/
public enum SemSegModelEnum {
DEEPLABV3("djl://ai.djl.pytorch/deeplabv3");
/**
* 根据名称获取枚举 (忽略大小写和下划线变体)
*/
public static SemSegModelEnum fromName(String name) {
String formatted = name.trim().toUpperCase().replaceAll("[-_]", "");
for (SemSegModelEnum model : values()) {
if (model.name().replaceAll("_", "").equals(formatted)) {
return model;
}
}
throw new IllegalArgumentException("未知模型名称: " + name);
}
private final String modelUri;
SemSegModelEnum(String modelUri) {
this.modelUri = modelUri;
}
public String getModelUri() {
return modelUri;
}
}

View File

@@ -0,0 +1,29 @@
package cn.smartjavaai.semseg.exception;
/**
* 语义分割异常
* @author dwj
*/
public class SemSegException extends RuntimeException{
public SemSegException() {
super();
}
public SemSegException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) {
super(message, cause, enableSuppression, writableStackTrace);
}
public SemSegException(String message, Throwable cause) {
super(message, cause);
}
public SemSegException(String message) {
super(message);
}
public SemSegException(Throwable cause) {
super(cause);
}
}

View File

@@ -0,0 +1,137 @@
package cn.smartjavaai.semseg.model;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.CategoryMask;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.Base64ImageUtils;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.semseg.config.SemSegModelConfig;
import cn.smartjavaai.semseg.criteria.SemSegCriteriaFactory;
import cn.smartjavaai.vision.utils.CategoryMaskFilter;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.impl.GenericObjectPool;
import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.util.Objects;
/**
* 语义分割模型
* @author dwj
*/
@Slf4j
public class CommonSemSegModel implements SemSegModel {
private SemSegModelConfig config;
private ZooModel<Image, CategoryMask> model;
private GenericObjectPool<Predictor<Image, CategoryMask>> predictorPool;
@Override
public void loadModel(SemSegModelConfig config) {
if(Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型枚举");
}
Criteria<Image, CategoryMask> criteria = SemSegCriteriaFactory.createCriteria(config);
this.config = config;
try {
model = criteria.loadModel();
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model));
int predictorPoolSize = config.getPredictorPoolSize();
if(config.getPredictorPoolSize() <= 0){
predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数
}
predictorPool.setMaxTotal(predictorPoolSize);
log.debug("当前设备: " + model.getNDManager().getDevice());
log.debug("当前引擎: " + Engine.getInstance().getEngineName());
log.debug("模型推理器线程池最大数量: " + predictorPoolSize);
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new DetectionException("模型加载失败", e);
}
}
@Override
public R<CategoryMask> detectBase64(String base64Image) {
if(StringUtils.isBlank(base64Image)){
return R.fail(R.Status.INVALID_IMAGE);
}
try {
byte[] imageData = Base64ImageUtils.base64ToImage(base64Image);
Image image = ImageFactory.getInstance().fromInputStream(new ByteArrayInputStream(imageData));
return detect(image);
} catch (IOException e) {
throw new DetectionException("读取图片异常", e);
}
}
@Override
public R<CategoryMask> detect(Image image) {
CategoryMask categoryMask = detectCore(image);
// 过滤
if(CollectionUtils.isNotEmpty(config.getAllowedClasses())
&& Objects.nonNull(categoryMask) && !categoryMask.getClasses().isEmpty()){
categoryMask = new CategoryMaskFilter(config.getAllowedClasses()).filter(categoryMask);
}
return R.ok(categoryMask);
}
/**
* 模型核心推理方法
* @param image
* @return
*/
public CategoryMask detectCore(Image image) {
Predictor<Image, CategoryMask> predictor = null;
try {
predictor = predictorPool.borrowObject();
return predictor.predict(image);
} catch (Exception e) {
throw new DetectionException("语义分割错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
log.debug("释放资源");
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public void close() throws Exception {
try {
if (predictorPool != null) {
predictorPool.close();
}
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (model != null) {
model.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,42 @@
package cn.smartjavaai.semseg.model;
import ai.djl.modality.Classifications;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.CategoryMask;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.semseg.config.SemSegModelConfig;
/**
* 语义分割模型
* @author dwj
*/
public interface SemSegModel extends AutoCloseable{
/**
* 加载模型
* @param config
*/
void loadModel(SemSegModelConfig config);
/**
* 语义分割
* @param base64Image
* @return
*/
default R<CategoryMask> detectBase64(String base64Image){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 语义分割
* @param image
* @return
*/
default R<CategoryMask> detect(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
}

View File

@@ -0,0 +1,102 @@
package cn.smartjavaai.vision.utils;
import ai.djl.modality.Classifications;
import ai.djl.modality.cv.output.CategoryMask;
import lombok.Data;
import java.util.ArrayList;
import java.util.List;
/**
* CategoryMask过滤器
* @author dwj
*/
@Data
public class CategoryMaskFilter {
private List<String> allowedClasses;
/**
* 构造函数,初始化允许的分类列表和置信度阈值
*
* @param allowedClasses 允许的分类列表null表示不过滤分类
*/
public CategoryMaskFilter(List<String> allowedClasses) {
this.allowedClasses = allowedClasses != null ? new ArrayList<>(allowedClasses) : null;
}
/**
* 对给定的CategoryMask进行过滤只保留允许的分类。
*
* @param categoryMask 待过滤的CategoryMask
* @return 过滤后的CategoryMask
*/
public CategoryMask filter(CategoryMask categoryMask) {
if (categoryMask == null) {
throw new NullPointerException("Input categoryMask cannot be null");
}
// If allowedClasses is null, return a copy of the original CategoryMask
if (allowedClasses == null) {
List<String> classNamesCopy = new ArrayList<>(categoryMask.getClasses());
int[][] maskCopy = copyMask(categoryMask.getMask());
return new CategoryMask(classNamesCopy, maskCopy);
}
// Get original class names and mask
List<String> originalClassNames = categoryMask.getClasses();
int[][] originalMask = categoryMask.getMask();
// Create a new class names list for allowed classes
List<String> filteredClassNames = new ArrayList<>();
// Map original indices to new indices for allowed classes
List<Integer> allowedOriginalIndices = new ArrayList<>();
for (int i = 0; i < originalClassNames.size(); i++) {
String className = originalClassNames.get(i);
if (allowedClasses.contains(className)) {
filteredClassNames.add(className);
allowedOriginalIndices.add(i);
}
}
// Create a new mask with the same dimensions
int height = originalMask.length;
int width = (height > 0) ? originalMask[0].length : 0;
int[][] filteredMask = new int[height][width];
// Update mask: remap pixels of allowed classes to new indices, others to 0
for (int h = 0; h < height; h++) {
for (int w = 0; w < width; w++) {
int originalIndex = originalMask[h][w];
int newIndex = allowedOriginalIndices.indexOf(originalIndex);
if (newIndex != -1) {
filteredMask[h][w] = newIndex;
} else {
filteredMask[h][w] = 0; // Background
}
}
}
return new CategoryMask(filteredClassNames, filteredMask);
}
/**
* Helper method to copy the mask array.
*
* @param originalMask the original mask to copy
* @return a deep copy of the mask
*/
private int[][] copyMask(int[][] originalMask) {
int height = originalMask.length;
int width = (height > 0) ? originalMask[0].length : 0;
int[][] copy = new int[height][width];
for (int h = 0; h < height; h++) {
System.arraycopy(originalMask[h], 0, copy[h], 0, width);
}
return copy;
}
}

View File

@@ -0,0 +1,62 @@
package cn.smartjavaai.vision.utils;
import ai.djl.modality.Classifications;
import lombok.Data;
import java.util.ArrayList;
import java.util.List;
import java.util.stream.Collectors;
/**
* Classification过滤器
* @author dwj
*/
@Data
public class ClassificationFilter {
private List<String> allowedClasses;
private float threshold;
/**
* 构造函数,初始化允许的分类列表和置信度阈值
*
* @param allowedClasses 允许的分类列表null表示不过滤分类
* @param threshold 置信度阈值
*/
public ClassificationFilter(List<String> allowedClasses, float threshold) {
this.allowedClasses = allowedClasses != null ? new ArrayList<>(allowedClasses) : null;
this.threshold = threshold;
}
/**
* 过滤分类结果并返回新的Classifications对象
*
* @param classifications 原始分类结果
* @return 过滤后的Classifications对象
*/
public Classifications filter(Classifications classifications) {
// 获取原始分类和概率
List<String> originalClasses = classifications.getClassNames();
List<Double> originalProbabilities = classifications.getProbabilities();
// 过滤分类和概率
List<String> filteredClasses = new ArrayList<>();
List<Double> filteredProbabilities = new ArrayList<>();
for (int i = 0; i < originalClasses.size(); i++) {
String className = originalClasses.get(i);
double probability = originalProbabilities.get(i);
// 检查是否满足置信度阈值和允许的分类列表
if (probability >= threshold && (allowedClasses == null || allowedClasses.contains(className))) {
filteredClasses.add(className);
filteredProbabilities.add(probability);
}
}
// 创建新的Classifications对象
return new Classifications(filteredClasses, filteredProbabilities);
}
}

View File

@@ -0,0 +1,84 @@
package cn.smartjavaai.vision.utils;
import ai.djl.modality.Classifications;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import lombok.Data;
import org.apache.commons.collections.CollectionUtils;
import java.util.ArrayList;
import java.util.Comparator;
import java.util.List;
import java.util.Objects;
/**
* Classification过滤器
* @author dwj
*/
@Data
public class DetectedObjectsFilter {
private List<String> allowedClasses;
private float threshold;
private int topK;
/**
* 构造函数,初始化允许的分类列表和置信度阈值
*
* @param allowedClasses 允许的分类列表null表示不过滤分类
* @param threshold 置信度阈值
*/
public DetectedObjectsFilter(List<String> allowedClasses, float threshold) {
this.allowedClasses = allowedClasses != null ? new ArrayList<>(allowedClasses) : null;
this.threshold = threshold;
}
public DetectedObjectsFilter(List<String> allowedClasses, float threshold, int topK) {
this.allowedClasses = allowedClasses;
this.threshold = threshold;
this.topK = topK;
}
/**
* 对给定的 DetectedObjects 进行过滤,返回过滤后的结果
*
* @param detectedObjects 待过滤的 DetectedObjects 对象
* @return 过滤后的 DetectedObjects 对象
*/
public DetectedObjects filter(DetectedObjects detectedObjects) {
if (Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0) {
return detectedObjects;
}
List<DetectedObjects.DetectedObject> items = detectedObjects.items();
// 按照允许的类别进行过滤
List<DetectedObjects.DetectedObject> filtered = new ArrayList<>();
//过滤类别
if(!CollectionUtils.isEmpty(allowedClasses)){
for (DetectedObjects.DetectedObject obj : items) {
if(allowedClasses.contains(obj.getClassName())){
filtered.add(obj);
}
}
}else{
filtered = items;
}
if(topK > 0 && filtered.size() > topK){
// 按照概率进行排序
filtered.sort(Comparator.comparingDouble(Classifications.Classification::getProbability).reversed());
filtered = filtered.subList(0, topK);
}
// 构建新的 DetectedObjects 返回
List<String> names = new ArrayList<>();
List<Double> probs = new ArrayList<>();
List<BoundingBox> boxes = new ArrayList<>();
for (DetectedObjects.DetectedObject obj : filtered) {
names.add(obj.getClassName());
probs.add(obj.getProbability());
boxes.add(obj.getBoundingBox());
}
return new DetectedObjects(names, probs, boxes);
}
}

View File

@@ -0,0 +1,114 @@
package cn.smartjavaai.vision.utils;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.output.Mask;
import ai.djl.modality.cv.output.Rectangle;
import cn.smartjavaai.common.entity.*;
import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.obb.entity.ObbResult;
import cn.smartjavaai.obb.entity.YoloRotatedBox;
import org.apache.commons.collections.CollectionUtils;
import org.opencv.core.Mat;
import org.opencv.core.Scalar;
import org.opencv.imgproc.Imgproc;
import javax.imageio.ImageIO;
import java.awt.image.BufferedImage;
import java.io.File;
import java.io.IOException;
import java.util.ArrayList;
import java.util.Iterator;
import java.util.List;
import java.util.Objects;
/**
* 目标检测相关工具类
* @author dwj
* @date 2025/4/9
*/
public class DetectorUtils {
/**
* 转换为FaceDetectedResult
* @param detection
* @param img
* @return
*/
public static DetectionResponse convertToDetectionResponse(DetectedObjects detection, Image img){
if(Objects.isNull(detection) || Objects.isNull(detection.getProbabilities())
|| detection.getProbabilities().isEmpty() || Objects.isNull(detection.items()) || detection.items().isEmpty()){
return null;
}
DetectionResponse detectionResponse = new DetectionResponse();
List<DetectionInfo> detectionInfoList = new ArrayList<DetectionInfo>();
List<DetectedObjects.DetectedObject> detectedObjectList = detection.items();
Iterator iterator = detectedObjectList.iterator();
int index = 0;
while(iterator.hasNext()) {
DetectedObjects.DetectedObject result = (DetectedObjects.DetectedObject)iterator.next();
String className = result.getClassName();
BoundingBox box = result.getBoundingBox();
int x = (int)(box.getBounds().getX() * (double)img.getWidth());
int y = (int)(box.getBounds().getY() * (double)img.getHeight());
int width = (int)(box.getBounds().getWidth() * (double)img.getWidth());
int height = (int)(box.getBounds().getHeight() * (double)img.getHeight());
DetectionRectangle rectangle = new DetectionRectangle(x, y, width, height);
DetectionInfo detectionInfo = new DetectionInfo(rectangle, detection.getProbabilities().get(index).floatValue());
//目标检测
if(box instanceof Rectangle){
ObjectDetInfo objectDetInfo = new ObjectDetInfo(className);
detectionInfo.setObjectDetInfo(objectDetInfo);
}else if(box instanceof Mask){
Mask mask = (Mask)box;
InstanceSegInfo instanceSegInfo = new InstanceSegInfo(className, mask.getProbDist());
detectionInfo.setInstanceSegInfo(instanceSegInfo);
}
detectionInfoList.add(detectionInfo);
index++;
}
detectionResponse.setDetectionInfoList(detectionInfoList);
return detectionResponse;
}
public static DetectionResponse obbToToDetectionResponse(ObbResult obbResult){
if(Objects.isNull(obbResult) || CollectionUtils.isEmpty(obbResult.getRotatedBoxeList())){
return null;
}
DetectionResponse detectionResponse = new DetectionResponse();
List<DetectionInfo> detectionInfoList = new ArrayList<DetectionInfo>();
for (YoloRotatedBox box : obbResult.getRotatedBoxeList()){
List<Point> points = box.toPoints();
RotatedBox rotatedBox = new RotatedBox(points.get(0), points.get(1), points.get(2), points.get(3));
ObbDetInfo obbDetInfo = new ObbDetInfo(box.className, rotatedBox);
DetectionInfo detectionInfo = new DetectionInfo();
detectionInfo.setScore(box.score);
detectionInfo.setObbDetInfo(obbDetInfo);
detectionInfoList.add(detectionInfo);
}
detectionResponse.setDetectionInfoList(detectionInfoList);
return detectionResponse;
}
/**
* 绘制文本框及文本
* @param srcMat
* @param rotatedBoxeList
*/
public static void drawRectWithText(Mat srcMat, List<YoloRotatedBox> rotatedBoxeList) {
for(YoloRotatedBox box : rotatedBoxeList){
List<Point> points = box.toPoints();
Imgproc.line(srcMat, points.get(0).toCvPoint(), points.get(1).toCvPoint(), new Scalar(0, 255, 0), 1);
Imgproc.line(srcMat, points.get(1).toCvPoint(), points.get(2).toCvPoint(), new Scalar(0, 255, 0),1);
Imgproc.line(srcMat, points.get(2).toCvPoint(), points.get(3).toCvPoint(), new Scalar(0, 255, 0),1);
Imgproc.line(srcMat, points.get(3).toCvPoint(), points.get(1).toCvPoint(), new Scalar(0, 255, 0), 1);
// 中文乱码
Imgproc.putText(srcMat, box.className, points.get(0).toCvPoint(), Imgproc.FONT_HERSHEY_SCRIPT_SIMPLEX, 1.0, new Scalar(0, 255, 0), 1);
}
}
}

View File

@@ -0,0 +1,45 @@
package cn.smartjavaai.vision.utils;
import cn.smartjavaai.obb.entity.ObbResult;
import cn.smartjavaai.obb.entity.YoloRotatedBox;
import java.util.Collections;
import java.util.List;
import java.util.stream.Collectors;
/**
* @author dwj
* @date 2025/8/24
*/
public class ObbResultFilter {
private List<String> allowedClasses;
private int topK;
public ObbResultFilter(List<String> allowedClasses, int topK) {
this.allowedClasses = allowedClasses;
this.topK = topK;
}
public ObbResult filter(ObbResult result) {
if (result == null || result.getRotatedBoxeList() == null) {
return new ObbResult(Collections.emptyList());
}
// 1. 按类别过滤
List<YoloRotatedBox> filtered = result.getRotatedBoxeList().stream()
.filter(box -> allowedClasses.contains(box.className))
.collect(Collectors.toList());
// 2. 按 score 排序
// filtered.sort((a, b) -> Float.compare(b.score, a.score));
// 3. 取前 topK
if (topK > 0 && filtered.size() > topK) {
filtered = filtered.subList(0, topK);
}
return new ObbResult(filtered);
}
}

View File

@@ -0,0 +1,78 @@
package cn.smartjavaai.vision.utils;
import ai.djl.util.JsonUtils;
import com.google.gson.annotations.SerializedName;
import org.apache.commons.io.FileUtils;
import java.io.BufferedInputStream;
import java.io.IOException;
import java.io.InputStream;
import java.net.URL;
import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.List;
import java.util.Map;
import java.util.Scanner;
import java.util.concurrent.ConcurrentHashMap;
/**
* Tensorflow synset utils
* @author dwj
*/
public class TensorflowSynsetUtils {
/**
* 加载synset
* @param synsetUrl
* @return
* @throws IOException
*/
public static Map<Integer, String> loadSynset(URL synsetUrl) throws IOException {
return loadSynset(synsetUrl.openStream());
}
/**
* 加载synset
* @param synsetPath
* @return
* @throws IOException
*/
public static Map<Integer, String> loadSynset(Path synsetPath) throws IOException {
return loadSynset(Files.newInputStream(synsetPath));
}
/**
* 加载synset
* @param inputStream
* @return
* @throws IOException
*/
public static Map<Integer, String> loadSynset(InputStream inputStream) throws IOException {
Map<Integer, String> map = new ConcurrentHashMap<>();
int maxId = 0;
try (InputStream is = new BufferedInputStream(inputStream);
Scanner scanner = new Scanner(is, StandardCharsets.UTF_8.name())) {
scanner.useDelimiter("item ");
while (scanner.hasNext()) {
String content = scanner.next();
content = content.replaceAll("(\"|\\d)\\n\\s", "$1,");
Item item = JsonUtils.GSON.fromJson(content, Item.class);
map.put(item.id, item.displayName);
if (item.id > maxId) {
maxId = item.id;
}
}
}
return map;
}
private static final class Item {
int id;
@SerializedName("display_name")
String displayName;
}
}