Skip to content
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
package org.jlab.rec.alert.AIPID;

import ai.djl.MalformedModelException;
import ai.djl.inference.Predictor;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.types.Shape;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import ai.djl.training.util.ProgressBar;
import ai.djl.translate.TranslateException;
import ai.djl.translate.Translator;
import ai.djl.translate.TranslatorContext;
import java.io.IOException;
import java.nio.file.Paths;
import java.util.logging.Logger;
import org.jlab.utils.CLASResources;

public class ModelPostPID {

private static final Logger LOGGER = Logger.getLogger(ModelPostPID.class.getName());
private static final int[] CLASS_IDS = {2212, 45, 46, 49, 47};
private static final int INPUT_SIZE = 18;

private final ZooModel<float[], float[]> model;

public ModelPostPID() {
System.setProperty("ai.djl.pytorch.num_interop_threads", "1");
System.setProperty("ai.djl.pytorch.num_threads", "1");
System.setProperty("ai.djl.pytorch.graph_optimizer", "false");

String path = CLASResources.getResourcePath("etc/data/nnet/rg-l/model_PID/");
Criteria<float[], float[]> criteria = Criteria.builder()
.setTypes(float[].class, float[].class)
.optModelPath(Paths.get(path))
.optEngine("PyTorch")
.optTranslator(translator())
.optProgress(new ProgressBar())
.build();
try {
model = criteria.loadModel();
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new RuntimeException(e);
}
}

public float[] prediction(float[] features) throws TranslateException {
if (features == null || features.length != INPUT_SIZE) {
LOGGER.warning("PostPID input must be float[18]");
return null;
}
try (Predictor<float[], float[]> predictor = model.newPredictor()) {
return predictor.predict(features);
}
}

private static Translator<float[], float[]> translator() {
return new Translator<>() {
@Override
public NDList processInput(TranslatorContext ctx, float[] features) {
return new NDList(ctx.getNDManager().create(features, new Shape(1, INPUT_SIZE)));
}

@Override
public float[] processOutput(TranslatorContext ctx, NDList output) {
float[] probabilities = output.get(0).toFloatArray();
int bestIndex = 0;
for (int i = 1; i < probabilities.length; i++) {
if (probabilities[i] > probabilities[bestIndex]) {
bestIndex = i;
}
}
return new float[]{
CLASS_IDS[bestIndex],
probabilities[0], probabilities[1], probabilities[2],
probabilities[3], probabilities[4]
};
}
};
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,7 @@

import ai.djl.MalformedModelException;
import ai.djl.inference.Predictor;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.ndarray.types.Shape;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
Expand All @@ -13,92 +11,96 @@
import ai.djl.translate.TranslateException;
import ai.djl.translate.Translator;
import ai.djl.translate.TranslatorContext;

import org.jlab.utils.CLASResources;

import java.io.IOException;
import java.nio.file.Paths;
import java.util.logging.Logger;
import org.jlab.utils.CLASResources;

public class ModelPrePID {

static final Logger LOGGER = Logger.getLogger(ModelPrePID.class.getName());
// Must match training class order
private static final int[] CLASS_IDS = new int[]{2212, 45, 46, 47, 49};

private final ZooModel<float[], float[]> model;

public ModelPrePID() {

Translator<float[], float[]> my_translator = new Translator<>() {
private static final Logger LOGGER = Logger.getLogger(ModelPrePID.class.getName());
private static final int[] CLASS_IDS = {2212, 45, 46, 49, 47};

@Override
public NDList processInput(TranslatorContext ctx, float[] floats) {
NDManager manager = ctx.getNDManager();
private final ZooModel<float[], float[]> ahdcModel;
private final ZooModel<float[], float[]> atofModel;

// IMPORTANT: model expects (batch, 23). Provide (1, 23).
NDArray x = manager.create(floats, new Shape(1, 23));
return new NDList(x);
}
public ModelPrePID() {
System.setProperty("ai.djl.pytorch.num_interop_threads", "1");
System.setProperty("ai.djl.pytorch.num_threads", "1");
System.setProperty("ai.djl.pytorch.graph_optimizer", "false");

@Override
public float[] processOutput(TranslatorContext ctx, NDList ndList) {
NDArray logits = ndList.get(0); // (1,5)
NDArray probs = logits.softmax(1); // (1,5)
ahdcModel = loadModel("model_prePID_AHDC", 11);
atofModel = loadModel("model_prePID_ATOF", 16);
}

float[] p = probs.toFloatArray(); // length 5 (row-major)
public ZooModel<float[], float[]> getModel() {
return ahdcModel;
}

// argmax
int bestIdx = 0;
float best = p[0];
for (int k = 1; k < 5; k++) {
if (p[k] > best) { best = p[k]; bestIdx = k; }
}
int prepid = CLASS_IDS[bestIdx];
public float[] prediction(float[] features) throws TranslateException {
if (features != null && features.length == 16) {
return predictionATOF(features);
}
return predictionAHDC(features);
}

// Return: prepid + probabilities in fixed class order
return new float[]{
(float) prepid,
p[0], p[1], p[2], p[3], p[4]
};
}
};
public float[] predictionAHDC(float[] features) throws TranslateException {
return predict(ahdcModel, features, 11);
}

System.setProperty("ai.djl.pytorch.num_interop_threads", "1");
System.setProperty("ai.djl.pytorch.num_threads", "1");
System.setProperty("ai.djl.pytorch.graph_optimizer", "false");
public float[] predictionATOF(float[] features) throws TranslateException {
return predict(atofModel, features, 16);
}

String path = CLASResources.getResourcePath("etc/data/nnet/rg-l/model_PrePID/");
private static float[] predict(ZooModel<float[], float[]> model, float[] features,
int expectedSize) throws TranslateException {
if (features == null || features.length != expectedSize) {
LOGGER.warning("PrePID input must be float[" + expectedSize + "]");
return null;
}
try (Predictor<float[], float[]> predictor = model.newPredictor()) {
return predictor.predict(features);
}
}

private static ZooModel<float[], float[]> loadModel(String directory, int inputSize) {
String path = CLASResources.getResourcePath("etc/data/nnet/rg-l/" + directory + "/");
Criteria<float[], float[]> criteria = Criteria.builder()
.setTypes(float[].class, float[].class)
.optModelPath(Paths.get(path))
.optEngine("PyTorch")
.optTranslator(my_translator)
.optTranslator(translator(inputSize))
.optProgress(new ProgressBar())
.build();

try {
model = criteria.loadModel();
return criteria.loadModel();
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new RuntimeException(e);
}
}

public ZooModel<float[], float[]> getModel() {
return model;
}
private static Translator<float[], float[]> translator(int inputSize) {
return new Translator<>() {
@Override
public NDList processInput(TranslatorContext ctx, float[] features) {
return new NDList(ctx.getNDManager().create(features, new Shape(1, inputSize)));
}

/** Returns float[]{prepid} where prepid in {2212,45,46,47,49}.
* @param features23
* @return
* @throws ai.djl.translate.TranslateException */
public float[] prediction(float[] features23) throws TranslateException {
if (features23 == null || features23.length != 23) {
LOGGER.warning("PrePID input must be float[23]");
return null;
}
Predictor<float[], float[]> predictor = model.newPredictor();
return predictor.predict(features23);
@Override
public float[] processOutput(TranslatorContext ctx, NDList output) {
float[] probabilities = output.get(0).toFloatArray();
int bestIndex = 0;
for (int i = 1; i < probabilities.length; i++) {
if (probabilities[i] > probabilities[bestIndex]) {
bestIndex = i;
}
}
return new float[]{
CLASS_IDS[bestIndex],
probabilities[0], probabilities[1], probabilities[2],
probabilities[3], probabilities[4]
};
}
};
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
package org.jlab.rec.alert.AIPID;

public class PIDResult {
public final int trackid;
public final int clusterid;
public final int pid;
public final float p2212, p45, p46, p47, p49;

public PIDResult(int trackid, int clusterid, float[] prediction) {
this.trackid = trackid;
this.clusterid = clusterid;
this.pid = (int) prediction[0];
this.p2212 = prediction[1];
this.p45 = prediction[2];
this.p46 = prediction[3];
this.p49 = prediction[4];
this.p47 = prediction[5];
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,8 @@ public class PrePIDResult {
public final int prepid;
public final float p2212, p45, p46, p47, p49;

public PrePIDResult(int trackid, int clusterid, int prepid, float p2212, float p45, float p46, float p47, float p49) {
public PrePIDResult(int trackid, int clusterid, int prepid,
float p2212, float p45, float p46, float p47, float p49) {
this.trackid = trackid;
this.clusterid = clusterid;
this.prepid = prepid;
Expand All @@ -16,4 +17,10 @@ public PrePIDResult(int trackid, int clusterid, int prepid, float p2212, float p
this.p47 = p47;
this.p49 = p49;
}

public PrePIDResult(int trackid, int clusterid, float[] prediction) {
this(trackid, clusterid, (int) prediction[0],
prediction[1], prediction[2], prediction[3],
prediction[5], prediction[4]);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,9 @@
import java.util.List;
import org.jlab.io.base.DataBank;
import org.jlab.io.base.DataEvent;
import org.jlab.rec.alert.AIPID.PIDResult;
import org.jlab.rec.alert.AIPID.PrePIDResult;
import org.jlab.rec.alert.projections.TrackProjection;
//import org.jlab.rec.alert.AIpid.PIDResult;

import ai.djl.util.Pair;

Expand All @@ -17,7 +18,7 @@
* @author Whit Armstrong
*/
public class RecoBankWriter {

/**
* Writes the bank of track projections.
*
Expand Down Expand Up @@ -54,12 +55,12 @@ public static DataBank fillProjectionsBank(DataEvent event, ArrayList<TrackProje
}
return bank;
}

/**
* Appends the alert match banks to an event.
*
* @param event the {@link DataEvent} in which to append the banks
* @param projections the {@link ArrayList} of {@link TrackProjection} containing the
* @param projections the {@link ArrayList} of {@link TrackProjection} containing the
* track projections info to be added
*
* @return 0 if it worked, 1 if it failed
Expand Down Expand Up @@ -91,17 +92,16 @@ public int appendTrackMatchingAIBank(DataEvent event, ArrayList<Pair<Integer, In

return 0;
}
public int appendPrePIDBank(DataEvent event, ArrayList<org.jlab.rec.alert.AIPID.PrePIDResult> results) {

public int appendPrePIDBank(DataEvent event, ArrayList<PrePIDResult> results) {

DataBank bank = event.createBank("ALERT::ai:prepid", results.size());
if (bank == null) {
System.err.println("COULD NOT CREATE A ALERT::ai:prepid BANK!!!!!!");
return 1;
}

for (int i = 0; i < results.size(); i++) {
org.jlab.rec.alert.AIPID.PrePIDResult r = results.get(i);
PrePIDResult r = results.get(i);
bank.setInt("trackid", i, r.trackid);
bank.setInt("clusterid", i, r.clusterid);
bank.setInt("prepid", i, r.prepid);
Expand All @@ -111,7 +111,27 @@ public int appendPrePIDBank(DataEvent event, ArrayList<org.jlab.rec.alert.AIPID.
bank.setFloat("p47", i, r.p47);
bank.setFloat("p49", i, r.p49);
}
event.appendBank(bank);
return 0;
}

public int appendPIDBank(DataEvent event, ArrayList<PIDResult> results) {
DataBank bank = event.createBank("ALERT::ai:pid", results.size());
if (bank == null) {
System.err.println("COULD NOT CREATE A ALERT::ai:pid BANK!!!!!!");
return 1;
}
for (int i = 0; i < results.size(); i++) {
PIDResult r = results.get(i);
bank.setInt("trackid", i, r.trackid);
bank.setInt("clusterid", i, r.clusterid);
bank.setInt("pid", i, r.pid);
bank.setFloat("prob_2212", i, r.p2212);
bank.setFloat("prob_45", i, r.p45);
bank.setFloat("prob_46", i, r.p46);
bank.setFloat("prob_47", i, r.p47);
bank.setFloat("prob_49", i, r.p49);
}
event.appendBank(bank);
return 0;
}
Expand Down
Loading
Loading