临时提交

This commit is contained in:
dengwenjie
2025-08-31 18:41:27 +08:00
parent 86ea7eb03e
commit 2b044fda29
25 changed files with 617 additions and 1005 deletions

View File

@@ -1,7 +1,11 @@
package cn.smartjavaai.common.utils;
import ai.djl.ndarray.NDArray;
import cn.smartjavaai.common.entity.R;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.Objects;
/**
* @author dwj
@@ -29,4 +33,14 @@ public class DJLCommonUtils {
return Files.exists(servingFile);
}
/**
* 判断 NDArray 是否为空
* @param ndArray
* @return
*/
public static boolean isNDArrayEmpty(NDArray ndArray){
return Objects.isNull(ndArray) || ndArray.size() == 0;
}
}

View File

@@ -1,5 +1,7 @@
package cn.smartjavaai.common.utils;
import ai.djl.modality.cv.output.Landmark;
import ai.djl.modality.cv.output.Point;
import ai.djl.modality.cv.output.Rectangle;
import ai.djl.modality.cv.util.NDImageUtils;
import ai.djl.ndarray.NDArray;
@@ -8,7 +10,9 @@ import ai.djl.ndarray.index.NDIndex;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
/**
* 按比例缩放,剩余空间用指定颜色填充
@@ -61,7 +65,7 @@ public class LetterBoxUtils {
// NDArray paddingImg = manager
// .full(new Shape(targetW, targetH, 3), padColor, DataType.UINT8);
NDArray paddingImg = manager.zeros(new Shape(targetW, targetH, 3), DataType.FLOAT32);
NDArray paddingImg = manager.zeros(new Shape(targetH, targetW, 3), DataType.FLOAT32);
paddingImg = paddingImg.add(114);
int padW = targetW - newW;
@@ -145,4 +149,65 @@ public class LetterBoxUtils {
return new Rectangle(x1, y1, boxW, boxH);
}
/**
* 恢复缩放后的 box(左上角坐标)
* @param landmark
* @param scale
* @param origImageWidth
* @param origImageHeight
*/
public static Landmark restoreBox(Landmark landmark, float scale, int origImageWidth, int origImageHeight, int inputWidth, int inputHeight, boolean isNormalized){
double x = 0;
double y = 0;
double width = 0;
double height = 0;
if(isNormalized){
x = landmark.getX() * inputWidth;
y = landmark.getY() * inputHeight;
width = landmark.getWidth() * inputWidth;
height = landmark.getHeight() * inputHeight;
}else{
x = landmark.getX();
y = landmark.getY();
width = landmark.getWidth();
height = landmark.getHeight();
}
double paddingWidth = (inputWidth - origImageWidth * scale) / 2;
double paddingHeight = (inputHeight - origImageHeight * scale) / 2;
// 去掉 padding
double x_noPad = x - paddingWidth;
double y_noPad = y - paddingHeight;
//模型输出就是原图坐标
double x1 = x_noPad / scale / origImageWidth;
double y1 = y_noPad / scale / origImageHeight;
double boxW = width / scale / origImageWidth ;
double boxH = height / scale / origImageHeight;
List<Point> points = new ArrayList<>();
// 要求关键点未归一化
landmark.getPath().forEach(point -> {
double pointX = (point.getX() - paddingWidth) / scale;
double pointY = (point.getY() - paddingHeight) / scale;
points.add(new Point(pointX, pointY));
});
return new Landmark(x1, y1, boxW, boxH, points);
}
/**
* 获取缩放后的图片大小
* @param origW 原始图片宽度
* @param origH 原始图片高度
* @param targetWidth 目标图片宽度
* @param targetHeight 目标图片高度
* @return
*/
public static int[] getResizeSize(int origW, int origH, int targetWidth, int targetHeight){
float r = Math.min(targetWidth / (float) origW, targetHeight / (float) origH);
int newW = Math.round(origW * r);
int newH = Math.round(origH * r);
return new int[]{newW, newH};
}
}

View File

@@ -81,16 +81,23 @@ public class NMSUtils {
*
*/
public static NDArray batchedNms(NDArray boxes, NDArray scores, NDArray idxs, float iouThreshold, NDManager manager) {
// System.out.println("---------------boxes:" + Arrays.toString(boxes.toFloatArray()));
List<NDArray> keepList = new ArrayList<>();
// 获取唯一 batch id
NDArray uniqueIdxs = idxs.unique().get(0);
for (long batchId : uniqueIdxs.toLongArray()) {
// 找出当前 batch 的框
NDArray mask = idxs.eq(batchId);
NDArray batchBoxes = boxes.get(mask);
NDArray batchScores = scores.get(mask);
// 执行单 batch NMS
int[] keepIndices = nms(batchBoxes, batchScores, iouThreshold);
int[] keepIndices = mtcnnNms(batchBoxes, batchScores, iouThreshold);
if (keepIndices.length > 0) {
// 将局部索引映射回全局索引
NDArray globalIndices = manager.arange(boxes.getShape().get(0))
@@ -101,10 +108,70 @@ public class NMSUtils {
keepList.add(globalIndices);
}
}
if (keepList.isEmpty()) {
return manager.create(new long[0]);
}
return NDArrays.concat(new NDList(keepList));
}
public static int[] mtcnnNms(NDArray boxes, NDArray scores, float iouThreshold) {
if (boxes.isEmpty()) {
return new int[0];
}
NDArray x1 = boxes.get(":, 0");
NDArray y1 = boxes.get(":, 1");
NDArray x2 = boxes.get(":, 2");
NDArray y2 = boxes.get(":, 3");
// 面积
NDArray areas = x2.sub(x1).add(1).mul(y2.sub(y1).add(1));
// scores 降序索引
NDArray order = scores.argSort();
//System.out.println("order" + order.getShape());
//System.out.println("order" + Arrays.toString(order.toLongArray()));
List<Integer> keep = new ArrayList<>();
while (order.size() > 0) {
int i = (int) order.getLong(-1);
keep.add(i);
if (order.size() == 1) break; // 没框了就退出
// 剩余框
NDArray idx = order.get("0:-1");
NDArray xx1 = x1.get(i).maximum(x1.get(idx));
NDArray yy1 = y1.get(i).maximum(y1.get(idx));
NDArray xx2 = x2.get(i).minimum(x2.get(idx));
NDArray yy2 = y2.get(i).minimum(y2.get(idx));
NDArray w = xx2.sub(xx1).add(1).maximum(0);
NDArray h = yy2.sub(yy1).add(1).maximum(0);
NDArray inter = w.mul(h);
NDArray union = areas.get(i).minimum(areas.get(idx));
NDArray iou = inter.div(union);
// System.out.println("Max IoU: " + iou.max().getFloat());
// System.out.println("Min IoU: " + iou.min().getFloat());
// System.out.println("Mean IoU: " + iou.mean().getFloat());
// System.out.println("Before: " + order.size());
// 保留 IoU <= 阈值的框
NDArray mask = iou.lte(iouThreshold);
// System.out.println("Mask size: " + mask.size() + " True count: " + mask.sum());
// 更新 order
order = idx.get(mask);
// System.out.println("After: " + order.size());
}
return keep.stream().mapToInt(Integer::intValue).toArray();
}
}