mirror of
https://github.com/geekwenjie/SmartJavaAI.git
synced 2026-09-10 03:28:49 +00:00
临时提交
This commit is contained in:
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
|
||||
@@ -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};
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user