Skip to content

Commit 77ab919

Browse files
authored
Merge pull request #42 from ajay1685/main
Fix ROI conversion from instance segmentation mask
2 parents 10869b4 + 9af2196 commit 77ab919

2 files changed

Lines changed: 54 additions & 50 deletions

File tree

build.gradle.kts

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@ plugins {
55

66
qupathExtension {
77
name = "qupath-extension-djl"
8-
version = "0.4.0"
8+
version = "0.4.1-SNAPSHOT"
99
group = "io.github.qupath"
1010
description = "QuPath extension to use Deep Java Library"
1111
automaticModule = "qupath.extension.djl"

src/main/java/qupath/ext/djl/DjlZoo.java

Lines changed: 53 additions & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -16,32 +16,6 @@
1616

1717
package qupath.ext.djl;
1818

19-
import java.awt.image.BandedSampleModel;
20-
import java.awt.image.BufferedImage;
21-
import java.awt.image.DataBufferFloat;
22-
import java.awt.image.WritableRaster;
23-
import java.io.IOException;
24-
import java.lang.reflect.Type;
25-
import java.net.URI;
26-
import java.util.ArrayList;
27-
import java.util.Arrays;
28-
import java.util.Collection;
29-
import java.util.Collections;
30-
import java.util.Comparator;
31-
import java.util.LinkedHashMap;
32-
import java.util.List;
33-
import java.util.Map;
34-
import java.util.Optional;
35-
import java.util.UUID;
36-
import java.util.concurrent.ConcurrentHashMap;
37-
import java.util.function.Function;
38-
import java.util.stream.Collectors;
39-
40-
import ai.djl.repository.MRL;
41-
import org.locationtech.jts.geom.util.AffineTransformation;
42-
import org.slf4j.Logger;
43-
import org.slf4j.LoggerFactory;
44-
4519
import ai.djl.Application;
4620
import ai.djl.MalformedModelException;
4721
import ai.djl.Model;
@@ -55,13 +29,14 @@
5529
import ai.djl.modality.cv.output.DetectedObjects;
5630
import ai.djl.modality.cv.output.DetectedObjects.DetectedObject;
5731
import ai.djl.modality.cv.output.Joints;
58-
import ai.djl.ndarray.NDList;
59-
import ai.djl.ndarray.types.LayoutType;
60-
import ai.djl.ndarray.types.Shape;
6132
import ai.djl.modality.cv.output.Landmark;
6233
import ai.djl.modality.cv.output.Mask;
6334
import ai.djl.modality.cv.translator.BigGANTranslator;
35+
import ai.djl.ndarray.NDList;
36+
import ai.djl.ndarray.types.LayoutType;
37+
import ai.djl.ndarray.types.Shape;
6438
import ai.djl.repository.Artifact;
39+
import ai.djl.repository.MRL;
6540
import ai.djl.repository.zoo.Criteria;
6641
import ai.djl.repository.zoo.ModelNotFoundException;
6742
import ai.djl.repository.zoo.ModelZoo;
@@ -72,6 +47,9 @@
7247
import ai.djl.util.ClassLoaderUtils;
7348
import ai.djl.util.Pair;
7449
import ai.djl.util.PairList;
50+
import org.locationtech.jts.geom.util.AffineTransformation;
51+
import org.slf4j.Logger;
52+
import org.slf4j.LoggerFactory;
7553
import qupath.lib.analysis.images.ContourTracing;
7654
import qupath.lib.analysis.images.SimpleImage;
7755
import qupath.lib.geom.Point2;
@@ -95,6 +73,27 @@
9573
import qupath.lib.roi.RoiTools;
9674
import qupath.lib.roi.interfaces.ROI;
9775

76+
import java.awt.image.BandedSampleModel;
77+
import java.awt.image.BufferedImage;
78+
import java.awt.image.DataBufferFloat;
79+
import java.awt.image.WritableRaster;
80+
import java.io.IOException;
81+
import java.lang.reflect.Type;
82+
import java.net.URI;
83+
import java.util.ArrayList;
84+
import java.util.Arrays;
85+
import java.util.Collection;
86+
import java.util.Collections;
87+
import java.util.Comparator;
88+
import java.util.LinkedHashMap;
89+
import java.util.List;
90+
import java.util.Map;
91+
import java.util.Optional;
92+
import java.util.UUID;
93+
import java.util.concurrent.ConcurrentHashMap;
94+
import java.util.function.Function;
95+
import java.util.stream.Collectors;
96+
9897
/**
9998
* Helper class for working with DeepJavaLibrary model zoos.
10099
*
@@ -316,13 +315,13 @@ private static TranslatorFactory getTranslatorFactory(String factoryClass) {
316315
* @return
317316
*/
318317
public static ROI createROI(DetectedObject obj, ImageRegion region) {
319-
var box = obj.getBoundingBox();
320-
if (box instanceof Mask) {
321-
return createROI((Mask)box, region, 0.5);
322-
} else if (box instanceof Landmark) {
323-
return createROI((Landmark)box, region);
318+
BoundingBox box = obj.getBoundingBox();
319+
if (box instanceof Mask mask) {
320+
return createROI(mask, region, 0.5);
321+
} else if (box instanceof Landmark landmark) {
322+
return createROI(landmark, region);
324323
} else
325-
return createROI((BoundingBox)box, region);
324+
return createROI(box, region);
326325
}
327326

328327
/**
@@ -335,17 +334,21 @@ public static ROI createROI(BoundingBox box, ImageRegion region) {
335334
var bounds = box.getBounds();
336335
double xo = 0.0;
337336
double yo = 0.0;
337+
double xScale = 1.0;
338+
double yScale = 1.0;
338339
var plane = ImagePlane.getDefaultPlane();
339340
if (region != null) {
340341
plane = region.getImagePlane();
341342
xo = region.getMinX();
342343
yo = region.getMinY();
344+
xScale = region.getWidth();
345+
yScale = region.getHeight();
343346
}
344347
return ROIs.createRectangleROI(
345-
xo + bounds.getX() * region.getWidth(),
346-
yo + bounds.getY() * region.getHeight(),
347-
bounds.getWidth() * region.getWidth(),
348-
bounds.getHeight() * region.getHeight(),
348+
xo + bounds.getX() * xScale,
349+
yo + bounds.getY() * yScale,
350+
bounds.getWidth() * xScale,
351+
bounds.getHeight() * yScale,
349352
plane);
350353
}
351354

@@ -358,17 +361,16 @@ public static ROI createROI(BoundingBox box, ImageRegion region) {
358361
*/
359362
public static ROI createROI(Mask mask, ImageRegion region, double threshold) {
360363
float[][] probs = mask.getProbDist();
361-
int w = probs.length;
362-
int h = probs[0].length;
364+
int h = probs.length;
365+
int w = probs[0].length;
363366
var buffer = new DataBufferFloat(w * h, 1);
364367
var sampleModel = new BandedSampleModel(buffer.getDataType(), w, h, 1);
365368
var raster = WritableRaster.createWritableRaster(sampleModel, buffer, null);
366-
for (int x = 0; x < w; x++) {
367-
float[] col = probs[x];
368-
for (int y = 0; y < h; y++) {
369-
raster.setSample(x, y, 0, col[y]);
370-
}
371-
}
369+
for (int y = 0; y < h; y++) {
370+
for (int x = 0; x < w; x++) {
371+
raster.setSample(x, y, 0, probs[y][x]);
372+
}
373+
}
372374
if (region == null)
373375
region = ImageRegion.createInstance(0, 0, w, h, 0, 0);
374376
var geometry = ContourTracing.createTracedGeometry(raster, threshold, Double.POSITIVE_INFINITY, 0, null);
@@ -377,8 +379,10 @@ public static ROI createROI(Mask mask, ImageRegion region, double threshold) {
377379

378380
var transform = new AffineTransformation();
379381
transform.scale(1.0/raster.getWidth(), 1.0/raster.getHeight());
380-
transform.scale(bounds.getWidth(), bounds.getHeight());
381-
transform.translate(bounds.getX(), bounds.getY());
382+
if(!mask.isFullImageMask()) {
383+
transform.scale(bounds.getWidth(), bounds.getHeight());
384+
transform.translate(bounds.getX(), bounds.getY());
385+
}
382386
transform.scale(region.getWidth(), region.getHeight());
383387
transform.translate(region.getX(), region.getY());
384388

0 commit comments

Comments
 (0)