Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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