diff --git a/src/main/java/org/beehive/gpullama3/inference/InferenceCoreBatchPrefillDecode.java b/src/main/java/org/beehive/gpullama3/inference/InferenceCoreBatchPrefillDecode.java
index c1458e42..6f78fd2d 100644
--- a/src/main/java/org/beehive/gpullama3/inference/InferenceCoreBatchPrefillDecode.java
+++ b/src/main/java/org/beehive/gpullama3/inference/InferenceCoreBatchPrefillDecode.java
@@ -1,6 +1,7 @@
package org.beehive.gpullama3.inference;
import org.beehive.gpullama3.auxiliary.Parallel;
+import org.beehive.gpullama3.inference.state.Qwen2MoEState;
import org.beehive.gpullama3.inference.state.State;
import org.beehive.gpullama3.inference.weights.standard.StandardWeights;
import org.beehive.gpullama3.inference.weights.tornado.TornadoWeights;
@@ -177,7 +178,7 @@ public static void batchForwardJavaPrefill(Model model, State state, int[] token
* then delegates graph execution to the plan.
*
* @param model
- * the LLaMA model
+ * the model
* @param state
* mutable inference state
* @param tokens
@@ -194,6 +195,9 @@ public static void batchForwardTornadoVMPrefill(Model model, State state, int[]
final TornadoWeights weights = (TornadoWeights) model.weights();
state.batchStartPosHolder.set(0, startPos);
+ if (state instanceof Qwen2MoEState moeState && moeState.activeBatchSizeHolder != null) {
+ moeState.activeBatchSizeHolder.set(0, chunkSize);
+ }
switch (weights.getWeightType()) {
case F16 -> {
@@ -230,7 +234,7 @@ public static void batchForwardTornadoVMPrefill(Model model, State state, int[]
* graph execution to the plan.
*
* @param model
- * the LLaMA model
+ * the model
* @param state
* mutable inference state
* @param token
diff --git a/src/main/java/org/beehive/gpullama3/inference/InferenceEngine.java b/src/main/java/org/beehive/gpullama3/inference/InferenceEngine.java
index 85ca2ec5..049d9b78 100644
--- a/src/main/java/org/beehive/gpullama3/inference/InferenceEngine.java
+++ b/src/main/java/org/beehive/gpullama3/inference/InferenceEngine.java
@@ -166,6 +166,7 @@ public static List generateTokensQwen3(Model model, State state, int st
// Storage for generated tokens
List generatedTokens = new ArrayList<>();
+ int generatedTokenBudget = Math.max(0, maxTokens - startPosition - promptTokens.size());
// Initialize token variables
int currentToken = state.latestToken; // BOS?
@@ -188,8 +189,11 @@ public static List generateTokensQwen3(Model model, State state, int st
if (echo) {
System.err.print(Tokenizer.replaceControlCharacters(model.tokenizer().decode(List.of(nextToken))));
}
- // We have reached the last prompt token and computed the first response-token.
- position++; // The current logit belongs to the next position
+ // The last prompt token produced the first response-token logits.
+ // The for-loop advances to the next sequence position.
+ if (generatedTokenBudget == 0) {
+ break;
+ }
} else {
// Mark the start of actual generation (after prompt processing)
if (inferenceStartNanos == 0) {
@@ -216,7 +220,7 @@ public static List generateTokensQwen3(Model model, State state, int st
}
// Check for stop condition
- if (stopTokens.contains(nextToken)) {
+ if (generatedTokens.size() >= generatedTokenBudget || stopTokens.contains(nextToken)) {
break;
}
@@ -393,6 +397,7 @@ public static List generateTokensGPUQwen3(Model model, State state, int
// prompt is longer than the token budget (actualMaxTokens), the difference is
// negative and would throw IllegalArgumentException("Illegal Capacity").
List generatedTokens = new ArrayList<>(Math.max(0, Math.min(256, actualMaxTokens - promptTokens.size()))); // Conservative estimate
+ int generatedTokenBudget = Math.max(0, actualMaxTokens - startPosition - promptTokens.size());
// Initialize token variables
int currentToken = state.latestToken; // BOS?
@@ -428,8 +433,11 @@ public static List generateTokensGPUQwen3(Model model, State state, int
if (echo) {
System.err.print(Tokenizer.replaceControlCharacters(model.tokenizer().decode(List.of(nextToken))));
}
- // We have reached the last prompt token and computed the first response-token.
- position++; // The current logit belongs to the next position
+ // The last prompt token produced the first response-token logits.
+ // The for-loop advances to the next sequence position.
+ if (generatedTokenBudget == 0) {
+ break;
+ }
} else {
// Mark the start of actual generation (after prompt processing)
if (inferenceStartNanos == 0) {
@@ -456,7 +464,7 @@ public static List generateTokensGPUQwen3(Model model, State state, int
}
// Check for stop condition
- if (stopTokens.contains(nextToken)) {
+ if (generatedTokens.size() >= generatedTokenBudget || stopTokens.contains(nextToken)) {
break;
}
@@ -678,4 +686,4 @@ public static List generateTokensGPUGranite(Model model, State state, i
return generatedTokens;
}
-}
\ No newline at end of file
+}
diff --git a/src/main/java/org/beehive/gpullama3/inference/InferenceEngineWithBatchPrefillDecode.java b/src/main/java/org/beehive/gpullama3/inference/InferenceEngineWithBatchPrefillDecode.java
index 4493340b..d811c214 100644
--- a/src/main/java/org/beehive/gpullama3/inference/InferenceEngineWithBatchPrefillDecode.java
+++ b/src/main/java/org/beehive/gpullama3/inference/InferenceEngineWithBatchPrefillDecode.java
@@ -147,9 +147,7 @@ public static List generateTokensLlama(Model model,
// @formatter:on
/**
- * LLaMA batched GPU prefill token generation (GPU, Phase 4).
- *
- * FP16 only; Q8_0 throws {@link UnsupportedOperationException}.
+ * Batched GPU prefill followed by single-token decode.
*
* Split loop:
*
@@ -187,17 +185,25 @@ public static List generateTokensGPULlama(Model model,
int N = promptTokens.size();
// ── Prefill ───────────────────────────────────────────────────────────
- // Build the token sequence at positions [startPosition .. startPosition+N-1]:
- // position startPosition+0 : currentToken (BOS/previous token)
- // position startPosition+k : promptTokens[k-1]
- int[] prefillSeq = new int[N];
- prefillSeq[0] = currentToken;
- for (int i = 1; i < N; i++) {
- prefillSeq[i] = promptTokens.get(i - 1);
+ // Qwen's regular path forwards promptTokens[0] directly at position 0.
+ // Keep the final prompt token for the B1 decode graph, which produces the
+ // first generation logits without duplicating the ChatML start token.
+ boolean qwen2MoE = model.getModelType() == org.beehive.gpullama3.model.ModelType.QWEN_2_MOE;
+ int prefillTokenCount = qwen2MoE ? Math.max(0, N - 1) : N;
+ int[] prefillSeq = new int[prefillTokenCount];
+ if (qwen2MoE) {
+ for (int i = 0; i < prefillTokenCount; i++) {
+ prefillSeq[i] = promptTokens.get(i);
+ }
+ } else {
+ prefillSeq[0] = currentToken;
+ for (int i = 1; i < N; i++) {
+ prefillSeq[i] = promptTokens.get(i - 1);
+ }
}
- for (int chunkStart = 0; chunkStart < N && pos + chunkStart < actualMaxTokens; chunkStart += batchSize) {
- int chunkEnd = Math.min(Math.min(chunkStart + batchSize, N), actualMaxTokens - pos);
+ for (int chunkStart = 0; chunkStart < prefillTokenCount && pos + chunkStart < actualMaxTokens; chunkStart += batchSize) {
+ int chunkEnd = Math.min(Math.min(chunkStart + batchSize, prefillTokenCount), actualMaxTokens - pos);
int chunkSize = chunkEnd - chunkStart;
int[] chunk = Arrays.copyOfRange(prefillSeq, chunkStart, chunkEnd);
@@ -213,12 +219,13 @@ public static List generateTokensGPULlama(Model model,
}
currentToken = promptTokens.get(N - 1);
- pos = startPosition + N;
+ pos = startPosition + (qwen2MoE ? N - 1 : N);
state.latestToken = currentToken;
long decodeStartNanos = System.nanoTime();
+ int generatedTokenBudget = Math.max(0, actualMaxTokens - startPosition - N);
// ── Decode ────────────────────────────────────────────────────────────
- while (pos < actualMaxTokens) {
+ while (pos < actualMaxTokens && generatedTokens.size() < generatedTokenBudget) {
var logits = InferenceCoreBatchPrefillDecode.forwardTornadoVMDecode(model, state, currentToken, pos, plan);
int nextToken = sampler.sampleToken(logits);
diff --git a/src/main/java/org/beehive/gpullama3/inference/state/Qwen2MoEState.java b/src/main/java/org/beehive/gpullama3/inference/state/Qwen2MoEState.java
index d7bdc92d..234978db 100644
--- a/src/main/java/org/beehive/gpullama3/inference/state/Qwen2MoEState.java
+++ b/src/main/java/org/beehive/gpullama3/inference/state/Qwen2MoEState.java
@@ -40,6 +40,19 @@ public class Qwen2MoEState extends Qwen2State {
public final FloatArray wrapSharedGate;
public final FloatArray wrapSharedOutput;
+ // TornadoVM buffers for the batch-prefill MoE path.
+ // Their shapes use the configured maximum batch size so TaskGraphs stay fixed.
+ public final FloatArray wrapRouterLogitsBatch;
+ public final IntArray activeBatchSizeHolder;
+ public final IntArray wrapSelectedExpertsBatch;
+ public final FloatArray wrapRoutingWeightsBatch;
+ public final IntArray wrapGroupedAssignmentIds;
+ public final IntArray wrapGroupedPositionByAssignment;
+ public final FloatArray wrapGroupedExpertHidden;
+ public final FloatArray wrapGroupedExpertDown;
+ public final FloatArray wrapSharedHiddenBatch;
+ public final FloatArray wrapSharedWeightBatch;
+
public Qwen2MoEState(Configuration config, int batchsize) {
super(config, batchsize);
Qwen2MoEConfiguration c = (Qwen2MoEConfiguration) config;
@@ -56,6 +69,33 @@ public Qwen2MoEState(Configuration config, int batchsize) {
this.wrapExpertGate = new FloatArray(c.moeHiddenDim() * c.numberOfExpertsUsed());
this.wrapSharedGate = new FloatArray(c.sharedExpertHiddenDim());
this.wrapSharedOutput = new FloatArray(c.dim());
+
+ int gpuBatchSize = Integer.getInteger("llama.prefillBatchSize", 1);
+ if (gpuBatchSize > 1) {
+ int assignments = gpuBatchSize * c.numberOfExpertsUsed();
+ this.wrapRouterLogitsBatch = new FloatArray(gpuBatchSize * c.numberOfExperts());
+ this.activeBatchSizeHolder = new IntArray(1);
+ this.activeBatchSizeHolder.init(gpuBatchSize);
+ this.wrapSelectedExpertsBatch = new IntArray(assignments);
+ this.wrapRoutingWeightsBatch = new FloatArray(assignments);
+ this.wrapGroupedAssignmentIds = new IntArray(assignments);
+ this.wrapGroupedPositionByAssignment = new IntArray(assignments);
+ this.wrapGroupedExpertHidden = new FloatArray(assignments * c.moeHiddenDim());
+ this.wrapGroupedExpertDown = new FloatArray(assignments * c.dim());
+ this.wrapSharedHiddenBatch = new FloatArray(gpuBatchSize * c.sharedExpertHiddenDim());
+ this.wrapSharedWeightBatch = new FloatArray(gpuBatchSize);
+ } else {
+ this.wrapRouterLogitsBatch = null;
+ this.activeBatchSizeHolder = null;
+ this.wrapSelectedExpertsBatch = null;
+ this.wrapRoutingWeightsBatch = null;
+ this.wrapGroupedAssignmentIds = null;
+ this.wrapGroupedPositionByAssignment = null;
+ this.wrapGroupedExpertHidden = null;
+ this.wrapGroupedExpertDown = null;
+ this.wrapSharedHiddenBatch = null;
+ this.wrapSharedWeightBatch = null;
+ }
}
@Override
diff --git a/src/main/java/org/beehive/gpullama3/model/qwen2/Qwen2MoE.java b/src/main/java/org/beehive/gpullama3/model/qwen2/Qwen2MoE.java
index 0a013e52..49967115 100644
--- a/src/main/java/org/beehive/gpullama3/model/qwen2/Qwen2MoE.java
+++ b/src/main/java/org/beehive/gpullama3/model/qwen2/Qwen2MoE.java
@@ -2,6 +2,7 @@
import org.beehive.gpullama3.inference.InferenceCore;
import org.beehive.gpullama3.inference.InferenceEngine;
+import org.beehive.gpullama3.inference.InferenceEngineWithBatchPrefillDecode;
import org.beehive.gpullama3.inference.sampler.Sampler;
import org.beehive.gpullama3.inference.state.Qwen2MoEState;
import org.beehive.gpullama3.inference.state.State;
@@ -90,7 +91,9 @@ public List generateTokens(State state, int startPosition, List generateTokensGPU(State state, int startPosition, List promptTokens, Set stopTokens, int maxTokens, Sampler sampler, boolean echo,
IntConsumer onTokenGenerated, TornadoVMMasterPlan tornadoVMPlan) {
if (WITH_PREFILL_DECODE && TornadoVMMasterPlan.PREFILL_BATCH_SIZE > 1) {
- throw new UnsupportedOperationException("Batch prefill/decode on GPU not yet implemented for Qwen2-MoE");
+ return InferenceEngineWithBatchPrefillDecode.generateTokensGPULlama(
+ this, state, startPosition, promptTokens, stopTokens, maxTokens,
+ sampler, echo, onTokenGenerated, tornadoVMPlan);
}
if (WITH_PREFILL_DECODE) {
throw new UnsupportedOperationException("Prefill/decode on GPU not yet implemented for Qwen2-MoE");
diff --git a/src/main/java/org/beehive/gpullama3/tornadovm/kernels/Qwen2MoEBatchKernels.java b/src/main/java/org/beehive/gpullama3/tornadovm/kernels/Qwen2MoEBatchKernels.java
new file mode 100644
index 00000000..ee685e7b
--- /dev/null
+++ b/src/main/java/org/beehive/gpullama3/tornadovm/kernels/Qwen2MoEBatchKernels.java
@@ -0,0 +1,619 @@
+package org.beehive.gpullama3.tornadovm.kernels;
+
+import uk.ac.manchester.tornado.api.KernelContext;
+import uk.ac.manchester.tornado.api.math.TornadoMath;
+import uk.ac.manchester.tornado.api.types.arrays.ByteArray;
+import uk.ac.manchester.tornado.api.types.arrays.FloatArray;
+import uk.ac.manchester.tornado.api.types.arrays.IntArray;
+
+/** GPU kernels used by Qwen2-MoE batch prefill. */
+public final class Qwen2MoEBatchKernels {
+
+ private static final int Q8_0_BLOCK_SIZE = 32;
+ private static final int Q8_0_BLOCK_BYTES = 34;
+
+ private Qwen2MoEBatchKernels() {}
+
+ /** Computes one router score for every token-expert pair. */
+ public static void batchedRouterProjection(
+ KernelContext context,
+ FloatArray input,
+ FloatArray routerLogits,
+ FloatArray routerWeights,
+ IntArray activeBatchSizeHolder,
+ int dim,
+ int numberOfExperts,
+ int localWorkGroupSize) {
+
+ int groupId = context.groupIdx;
+ int localId = context.localIdx;
+
+ // One work-group handles one (token, expert) pair.
+ // For 60 experts, groups 0..59 process token 0, groups 60..119
+ // process token 1, and so on.
+
+ int token = groupId / numberOfExperts;
+ int expert = groupId % numberOfExperts;
+ if (token >= activeBatchSizeHolder.get(0)) {
+ return;
+ }
+
+ int inputOffset = token * dim;
+ int weightOffset = expert * dim;
+
+ // Each local thread computes a strided part of the dot product.
+ float partialSum = 0.0f;
+
+ for (int column = localId; column < dim; column += localWorkGroupSize) {
+
+ partialSum +=
+ input.get(inputOffset + column) * routerWeights.get(weightOffset + column);
+ }
+
+ // Store every thread's partial sum in local memory.
+ float[] localSums = context.allocateFloatLocalArray(localWorkGroupSize);
+
+ localSums[localId] = partialSum;
+ context.localBarrier();
+
+ // Reduce the partial sums inside this work-group.
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+
+ context.localBarrier();
+ }
+
+ // Local thread 0 writes routerLogits[token, expert].
+ if (localId == 0) {
+ int outputOffset = token * numberOfExperts + expert;
+ routerLogits.set(outputOffset, localSums[0]);
+ }
+ }
+
+ /** Adds Qwen2's Q, K, and V biases independently for every active token. */
+ public static void batchedQKVBias(
+ KernelContext context,
+ FloatArray qBatch,
+ FloatArray kBatch,
+ FloatArray vBatch,
+ FloatArray qBias,
+ FloatArray kBias,
+ FloatArray vBias,
+ IntArray activeBatchSizeHolder,
+ int dim,
+ int kvDim) {
+
+ int index = context.globalIdx;
+ int rowsPerToken = dim + 2 * kvDim;
+ int token = index / rowsPerToken;
+ int row = index % rowsPerToken;
+
+ if (token >= activeBatchSizeHolder.get(0)) {
+ return;
+ }
+
+ if (row < dim) {
+ int qIndex = token * dim + row;
+ qBatch.set(qIndex, qBatch.get(qIndex) + qBias.get(row));
+ } else if (row < dim + kvDim) {
+ int kRow = row - dim;
+ int kIndex = token * kvDim + kRow;
+ kBatch.set(kIndex, kBatch.get(kIndex) + kBias.get(kRow));
+ } else {
+ int vRow = row - dim - kvDim;
+ int vIndex = token * kvDim + vRow;
+ vBatch.set(vIndex, vBatch.get(vIndex) + vBias.get(vRow));
+ }
+ }
+
+ /** Applies Qwen2 split-half RoPE and writes active K/V rows to the cache. */
+ public static void batchedRopeWithKVCacheQwen2(
+ KernelContext context,
+ IntArray batchStartPosHolder,
+ IntArray activeBatchSizeHolder,
+ FloatArray qBatch,
+ FloatArray kBatch,
+ FloatArray vBatch,
+ FloatArray keyCache,
+ FloatArray valueCache,
+ int kvDim,
+ int headSize,
+ int layerIndex,
+ int contextLength,
+ int dim,
+ float ropeTheta) {
+
+ int index = context.globalIdx;
+ int pairsPerToken = dim / 2;
+ int token = index / pairsPerToken;
+ int pair = index % pairsPerToken;
+
+ if (token >= activeBatchSizeHolder.get(0)) {
+ return;
+ }
+
+ int position = batchStartPosHolder.get(0) + token;
+ int halfHeadSize = headSize / 2;
+ int component = pair % halfHeadSize;
+ int head = pair / halfHeadSize;
+
+ float frequency = 1.0f / TornadoMath.pow(ropeTheta, 2.0f * component / (float) headSize);
+ float angle = position * frequency;
+ float cosine = TornadoMath.cos(angle);
+ float sine = TornadoMath.sin(angle);
+
+ int qHeadOffset = token * dim + head * headSize;
+ float q0 = qBatch.get(qHeadOffset + component);
+ float q1 = qBatch.get(qHeadOffset + component + halfHeadSize);
+ qBatch.set(qHeadOffset + component, q0 * cosine - q1 * sine);
+ qBatch.set(qHeadOffset + component + halfHeadSize, q0 * sine + q1 * cosine);
+
+ if (pair < kvDim / 2) {
+ int kvHead = pair / halfHeadSize;
+ int kHeadOffset = token * kvDim + kvHead * headSize;
+ float k0 = kBatch.get(kHeadOffset + component);
+ float k1 = kBatch.get(kHeadOffset + component + halfHeadSize);
+ float rotatedK0 = k0 * cosine - k1 * sine;
+ float rotatedK1 = k0 * sine + k1 * cosine;
+
+ kBatch.set(kHeadOffset + component, rotatedK0);
+ kBatch.set(kHeadOffset + component + halfHeadSize, rotatedK1);
+
+ int cacheOffset =
+ layerIndex * contextLength * kvDim + position * kvDim + kvHead * headSize;
+ keyCache.set(cacheOffset + component, rotatedK0);
+ keyCache.set(cacheOffset + component + halfHeadSize, rotatedK1);
+ valueCache.set(cacheOffset + component, vBatch.get(kHeadOffset + component));
+ valueCache.set(
+ cacheOffset + component + halfHeadSize,
+ vBatch.get(kHeadOffset + component + halfHeadSize));
+ }
+ }
+
+ /** Applies softmax and selects Top-K experts independently for each token. */
+ public static void batchedSoftmaxAndTopK(
+ KernelContext context,
+ FloatArray routerLogits,
+ IntArray selectedExperts,
+ FloatArray routingWeights,
+ IntArray activeBatchSizeHolder,
+ int numberOfExperts,
+ int topK) {
+
+ // One GPU thread handles the complete routing result for one token.
+ int token = context.globalIdx;
+ if (token >= activeBatchSizeHolder.get(0)) {
+ return;
+ }
+
+ int logitsOffset = token * numberOfExperts;
+ int assignmentOffset = token * topK;
+
+ // Find this token's maximum router logit.
+ float maxLogit = Float.NEGATIVE_INFINITY;
+ for (int i = logitsOffset; i < logitsOffset + numberOfExperts; i++) {
+ maxLogit = Math.max(maxLogit, routerLogits.get(i));
+ }
+
+ // Compute the stable softmax denominator.
+ float sumExp = 0.0f;
+
+ for (int expert = 0; expert < numberOfExperts; expert++) {
+ float logit = routerLogits.get(logitsOffset + expert);
+ sumExp += TornadoMath.exp(logit - maxLogit);
+ }
+
+ // Convert this token's logits to probabilities.
+ for (int expert = 0; expert < numberOfExperts; expert++) {
+ int index = logitsOffset + expert;
+ float logit = routerLogits.get(index);
+
+ float probability = TornadoMath.exp(logit - maxLogit) / sumExp;
+
+ routerLogits.set(index, probability);
+ }
+
+ // Select Top-K expert IDs and their routing weights.
+ for (int slot = 0; slot < topK; slot++) {
+ float currentMax = Float.NEGATIVE_INFINITY;
+ int selectedExpert = -1;
+
+ for (int expert = 0; expert < numberOfExperts; expert++) {
+ float probability = routerLogits.get(logitsOffset + expert);
+
+ if (probability > currentMax) {
+ currentMax = probability;
+ selectedExpert = expert;
+ }
+ }
+
+ selectedExperts.set(assignmentOffset + slot, selectedExpert);
+ routingWeights.set(assignmentOffset + slot, currentMax);
+
+ routerLogits.set(logitsOffset + selectedExpert, Float.NEGATIVE_INFINITY);
+ }
+ }
+
+ /** Groups token-expert assignments by expert. */
+ public static void groupAssignmentsByExpert(
+ KernelContext context,
+ IntArray selectedExperts,
+ IntArray groupedAssignmentIds,
+ IntArray groupedPositionByAssignment,
+ IntArray activeBatchSizeHolder,
+ int numberOfExperts,
+ int topK) {
+
+ // Start with one thread for a simple, deterministic implementation.
+ if (context.globalIdx != 0) {
+ return;
+ }
+
+ int numberOfAssignments = activeBatchSizeHolder.get(0) * topK;
+ int groupedPosition = 0;
+
+ for (int expert = 0; expert < numberOfExperts; expert++) {
+ for (int assignment = 0; assignment < numberOfAssignments; assignment++) {
+ int selectedExpert = selectedExperts.get(assignment);
+ if (selectedExpert == expert) {
+ groupedAssignmentIds.set(groupedPosition, assignment);
+ groupedPositionByAssignment.set(assignment, groupedPosition);
+ groupedPosition++;
+ }
+ }
+ }
+ }
+
+ /** Computes routed Gate/Up projections in expert-grouped assignment order. */
+ public static void groupedRoutedExpertsGateUpSwiGLUQ8_0(
+ KernelContext context,
+ FloatArray inputBatch,
+ IntArray selectedExperts,
+ IntArray groupedAssignmentIds,
+ IntArray activeBatchSizeHolder,
+ ByteArray gateExperts,
+ ByteArray upExperts,
+ FloatArray groupedExpertHidden,
+ int dim,
+ int moeHiddenDim,
+ int numberOfExperts,
+ int topK,
+ int localWorkGroupSize) {
+
+ int flatGroupId = context.groupIdx;
+ int localId = context.localIdx;
+
+ int groupedPosition = flatGroupId / moeHiddenDim;
+ int rowId = flatGroupId % moeHiddenDim;
+ int numberOfAssignments = activeBatchSizeHolder.get(0) * topK;
+
+ boolean active = groupedPosition < numberOfAssignments;
+ int assignment = 0;
+ int token = 0;
+ int expert = 0;
+ if (active) {
+ assignment = groupedAssignmentIds.get(groupedPosition);
+ token = assignment / topK;
+ expert = selectedExperts.get(assignment);
+ active = expert >= 0 && expert < numberOfExperts;
+ }
+
+ int blocksPerRow = (dim + Q8_0_BLOCK_SIZE - 1) / Q8_0_BLOCK_SIZE;
+ int rowBlockOffset = (expert * moeHiddenDim + rowId) * blocksPerRow;
+ int inputOffset = token * dim;
+
+ float gatePartialSum = 0.0f;
+ float upPartialSum = 0.0f;
+ if (active) {
+ for (int column = localId; column < dim; column += localWorkGroupSize) {
+
+ int blockByteOffset =
+ (rowBlockOffset + column / Q8_0_BLOCK_SIZE) * Q8_0_BLOCK_BYTES;
+ int quantOffset = blockByteOffset + 2 + column % Q8_0_BLOCK_SIZE;
+
+ float inputValue = inputBatch.get(inputOffset + column);
+ float gateScale = gateExperts.getHalfFloat(blockByteOffset).getFloat32();
+ float upScale = upExperts.getHalfFloat(blockByteOffset).getFloat32();
+
+ gatePartialSum += gateExperts.get(quantOffset) * gateScale * inputValue;
+ upPartialSum += upExperts.get(quantOffset) * upScale * inputValue;
+ }
+ }
+
+ float[] localSums = context.allocateFloatLocalArray(localWorkGroupSize);
+
+ localSums[localId] = gatePartialSum;
+ context.localBarrier();
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+ context.localBarrier();
+ }
+ float gate = localSums[0];
+
+ localSums[localId] = upPartialSum;
+ context.localBarrier();
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+ context.localBarrier();
+ }
+
+ if (localId == 0 && active) {
+ float up = localSums[0];
+ float siluGate = gate / (1.0f + TornadoMath.exp(-gate));
+ int outputOffset = groupedPosition * moeHiddenDim + rowId;
+ groupedExpertHidden.set(outputOffset, siluGate * up);
+ }
+ }
+
+ /** Down-projects every routed assignment in expert-grouped order. */
+ public static void groupedRoutedExpertsDownQ8_0(
+ KernelContext context,
+ FloatArray groupedExpertHidden,
+ IntArray selectedExperts,
+ IntArray groupedAssignmentIds,
+ IntArray activeBatchSizeHolder,
+ ByteArray downExperts,
+ FloatArray groupedExpertDown,
+ int dim,
+ int moeHiddenDim,
+ int numberOfExperts,
+ int topK,
+ int localWorkGroupSize) {
+
+ int flatGroupId = context.groupIdx;
+ int localId = context.localIdx;
+
+ int groupedPosition = flatGroupId / dim;
+ int rowId = flatGroupId % dim;
+ int numberOfAssignments = activeBatchSizeHolder.get(0) * topK;
+
+ boolean active = groupedPosition < numberOfAssignments;
+ int assignment = 0;
+ int expert = 0;
+ if (active) {
+ assignment = groupedAssignmentIds.get(groupedPosition);
+ expert = selectedExperts.get(assignment);
+ active = expert >= 0 && expert < numberOfExperts;
+ }
+
+ int blocksPerRow = (moeHiddenDim + Q8_0_BLOCK_SIZE - 1) / Q8_0_BLOCK_SIZE;
+ int rowBlockOffset = (expert * dim + rowId) * blocksPerRow;
+ int hiddenOffset = groupedPosition * moeHiddenDim;
+
+ float partialSum = 0.0f;
+ if (active) {
+ for (int column = localId; column < moeHiddenDim; column += localWorkGroupSize) {
+
+ int blockByteOffset =
+ (rowBlockOffset + column / Q8_0_BLOCK_SIZE) * Q8_0_BLOCK_BYTES;
+ int quantOffset = blockByteOffset + 2 + column % Q8_0_BLOCK_SIZE;
+
+ float weight =
+ downExperts.get(quantOffset)
+ * downExperts.getHalfFloat(blockByteOffset).getFloat32();
+ partialSum += weight * groupedExpertHidden.get(hiddenOffset + column);
+ }
+ }
+
+ float[] localSums = context.allocateFloatLocalArray(localWorkGroupSize);
+ localSums[localId] = partialSum;
+ context.localBarrier();
+
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+ context.localBarrier();
+ }
+
+ if (localId == 0 && active) {
+ int outputOffset = groupedPosition * dim + rowId;
+ groupedExpertDown.set(outputOffset, localSums[0]);
+ }
+ }
+
+ /** Adds the weighted routed-expert results back to each token's residual. */
+ public static void accumulateGroupedRoutedExperts(
+ KernelContext context,
+ FloatArray groupedExpertDown,
+ IntArray groupedPositionByAssignment,
+ FloatArray routingWeights,
+ FloatArray residualBatch,
+ IntArray activeBatchSizeHolder,
+ int dim,
+ int topK) {
+
+ int index = context.globalIdx;
+ int token = index / dim;
+ int rowId = index % dim;
+
+ if (token >= activeBatchSizeHolder.get(0)) {
+ return;
+ }
+
+ float result = residualBatch.get(index);
+ int assignmentOffset = token * topK;
+ for (int slot = 0; slot < topK; slot++) {
+ int assignment = assignmentOffset + slot;
+ int groupedPosition = groupedPositionByAssignment.get(assignment);
+ int downOffset = groupedPosition * dim + rowId;
+ result += routingWeights.get(assignment) * groupedExpertDown.get(downOffset);
+ }
+
+ residualBatch.set(index, result);
+ }
+
+ /** Computes the shared expert Gate/Up result independently for each token. */
+ public static void batchedSharedExpertGateUpSwiGLUQ8_0(
+ KernelContext context,
+ FloatArray inputBatch,
+ IntArray activeBatchSizeHolder,
+ ByteArray sharedGate,
+ ByteArray sharedUp,
+ FloatArray sharedHiddenBatch,
+ int dim,
+ int sharedExpertHiddenDim,
+ int localWorkGroupSize) {
+
+ int flatGroupId = context.groupIdx;
+ int localId = context.localIdx;
+
+ int token = flatGroupId / sharedExpertHiddenDim;
+ int rowId = flatGroupId % sharedExpertHiddenDim;
+ boolean active = token < activeBatchSizeHolder.get(0);
+
+ int blocksPerRow = (dim + Q8_0_BLOCK_SIZE - 1) / Q8_0_BLOCK_SIZE;
+ int rowBlockOffset = rowId * blocksPerRow;
+ int inputOffset = token * dim;
+
+ float gatePartialSum = 0.0f;
+ float upPartialSum = 0.0f;
+ if (active) {
+ for (int column = localId; column < dim; column += localWorkGroupSize) {
+
+ int blockByteOffset =
+ (rowBlockOffset + column / Q8_0_BLOCK_SIZE) * Q8_0_BLOCK_BYTES;
+ int quantOffset = blockByteOffset + 2 + column % Q8_0_BLOCK_SIZE;
+
+ float inputValue = inputBatch.get(inputOffset + column);
+ float gateScale = sharedGate.getHalfFloat(blockByteOffset).getFloat32();
+ float upScale = sharedUp.getHalfFloat(blockByteOffset).getFloat32();
+
+ gatePartialSum += sharedGate.get(quantOffset) * gateScale * inputValue;
+ upPartialSum += sharedUp.get(quantOffset) * upScale * inputValue;
+ }
+ }
+
+ float[] localSums = context.allocateFloatLocalArray(localWorkGroupSize);
+
+ localSums[localId] = gatePartialSum;
+ context.localBarrier();
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+ context.localBarrier();
+ }
+ float gate = localSums[0];
+
+ localSums[localId] = upPartialSum;
+ context.localBarrier();
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+ context.localBarrier();
+ }
+
+ if (localId == 0 && active) {
+ float up = localSums[0];
+ float siluGate = gate / (1.0f + TornadoMath.exp(-gate));
+ int outputOffset = token * sharedExpertHiddenDim + rowId;
+ sharedHiddenBatch.set(outputOffset, siluGate * up);
+ }
+ }
+
+ /** Computes the sigmoid gate that scales each token's shared-expert output. */
+ public static void batchedSharedExpertGateWeight(
+ KernelContext context,
+ FloatArray inputBatch,
+ FloatArray sharedGateInput,
+ FloatArray sharedWeightBatch,
+ IntArray activeBatchSizeHolder,
+ int dim,
+ int localWorkGroupSize) {
+
+ int token = context.groupIdx;
+ int localId = context.localIdx;
+ boolean active = token < activeBatchSizeHolder.get(0);
+ int inputOffset = token * dim;
+
+ float partialScore = 0.0f;
+ if (active) {
+ for (int column = localId; column < dim; column += localWorkGroupSize) {
+ partialScore += sharedGateInput.get(column) * inputBatch.get(inputOffset + column);
+ }
+ }
+
+ float[] localSums = context.allocateFloatLocalArray(localWorkGroupSize);
+ localSums[localId] = partialScore;
+ context.localBarrier();
+
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+ context.localBarrier();
+ }
+
+ if (localId == 0 && active) {
+ float sharedWeight = 1.0f / (1.0f + TornadoMath.exp(-localSums[0]));
+ sharedWeightBatch.set(token, sharedWeight);
+ }
+ }
+
+ /** Down-projects the shared expert and adds it to each token's residual. */
+ public static void batchedSharedExpertDownAndAccumulateQ8_0(
+ KernelContext context,
+ FloatArray sharedHiddenBatch,
+ FloatArray sharedWeightBatch,
+ IntArray activeBatchSizeHolder,
+ ByteArray sharedDown,
+ FloatArray residualBatch,
+ int dim,
+ int sharedExpertHiddenDim,
+ int localWorkGroupSize) {
+
+ int flatGroupId = context.groupIdx;
+ int localId = context.localIdx;
+
+ int token = flatGroupId / dim;
+ int rowId = flatGroupId % dim;
+ boolean active = token < activeBatchSizeHolder.get(0);
+
+ int blocksPerRow = (sharedExpertHiddenDim + Q8_0_BLOCK_SIZE - 1) / Q8_0_BLOCK_SIZE;
+ int rowBlockOffset = rowId * blocksPerRow;
+ int hiddenOffset = token * sharedExpertHiddenDim;
+
+ float partialSum = 0.0f;
+ if (active) {
+ for (int column = localId;
+ column < sharedExpertHiddenDim;
+ column += localWorkGroupSize) {
+
+ int blockByteOffset =
+ (rowBlockOffset + column / Q8_0_BLOCK_SIZE) * Q8_0_BLOCK_BYTES;
+ int quantOffset = blockByteOffset + 2 + column % Q8_0_BLOCK_SIZE;
+
+ float weight =
+ sharedDown.get(quantOffset)
+ * sharedDown.getHalfFloat(blockByteOffset).getFloat32();
+ partialSum += weight * sharedHiddenBatch.get(hiddenOffset + column);
+ }
+ }
+
+ float[] localSums = context.allocateFloatLocalArray(localWorkGroupSize);
+ localSums[localId] = partialSum;
+ context.localBarrier();
+
+ for (int stride = localWorkGroupSize / 2; stride > 0; stride >>= 1) {
+ if (localId < stride) {
+ localSums[localId] += localSums[localId + stride];
+ }
+ context.localBarrier();
+ }
+
+ if (localId == 0 && active) {
+ int outputOffset = token * dim + rowId;
+ float weightedOutput = sharedWeightBatch.get(token) * localSums[0];
+ residualBatch.set(outputOffset, residualBatch.get(outputOffset) + weightedOutput);
+ }
+ }
+}
diff --git a/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/Qwen2MoEQ8_0FFNLayers.java b/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/Qwen2MoEQ8_0FFNLayers.java
index ee7a07e0..9b7a250d 100644
--- a/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/Qwen2MoEQ8_0FFNLayers.java
+++ b/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/Qwen2MoEQ8_0FFNLayers.java
@@ -24,10 +24,10 @@
* routed-expert pipeline: normalize, route, choose top-K experts, execute each
* selected expert, and accumulate its weighted output into {@code wrapX}.
*/
-public final class Qwen2MoEQ8_0FFNLayers
+public class Qwen2MoEQ8_0FFNLayers
extends AbstractTransformerLayerTaskGraphs {
- private final Qwen2MoEState moeState;
+ protected final Qwen2MoEState moeState;
public Qwen2MoEQ8_0FFNLayers(String taskGraphName,
Qwen2MoEState state,
@@ -105,9 +105,29 @@ private WorkerGrid workerForRows(int rows) {
protected TaskGraph createFFNLayerTaskGraph(int layerIndex) {
TaskGraph layer = new TaskGraph("layer_" + layerIndex);
// Reuse wrapX produced by the previous TaskGraph on the GPU.
- layer.consumeFromDevice(moeState.wrapX);
- // Upload this layer's read-only weights from CPU to GPU on the first execution.
- layer.transferToDevice(DataTransferMode.FIRST_EXECUTION,
+ String predecessor = predecessorGraphName(layerIndex);
+ if (predecessor == null) {
+ layer.consumeFromDevice(moeState.wrapX);
+ } else {
+ layer.consumeFromDevice(predecessor, moeState.wrapX);
+ }
+ layer = configureLayerWeights(layer, layerIndex);
+ layer = configureLayerDataTransfers(layer, layerIndex);
+
+ configureAttention(layer, layerIndex);
+ configureRoutedExperts(layer, layerIndex);
+ layer.persistOnDevice(moeState.wrapX, moeState.wrapKeyCache, moeState.wrapValueCache);
+ return layer;
+ }
+
+ /** Returns an explicit predecessor for plans that connect multiple graph chains. */
+ protected String predecessorGraphName(int layerIndex) {
+ return null;
+ }
+
+ /** Uploads this layer's read-only weights on its first execution. */
+ protected TaskGraph configureLayerWeights(TaskGraph layer, int layerIndex) {
+ return layer.transferToDevice(DataTransferMode.FIRST_EXECUTION,
weights.rms_att_weightLayered[layerIndex].asFloatArray(),
weights.wqLayered[layerIndex].asByteArray(),
weights.wkLayered[layerIndex].asByteArray(),
@@ -125,12 +145,6 @@ protected TaskGraph createFFNLayerTaskGraph(int layerIndex) {
weights.sharedUpLayered[layerIndex].asByteArray(),
weights.sharedDownLayered[layerIndex].asByteArray(),
weights.sharedGateInputLayered[layerIndex].asFloatArray());
- layer = configureLayerDataTransfers(layer, layerIndex);
-
- configureAttention(layer, layerIndex);
- configureRoutedExperts(layer, layerIndex);
- layer.persistOnDevice(moeState.wrapX);
- return layer;
}
/** Adds the normal Qwen2 attention tasks to this layer's TaskGraph. */
diff --git a/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/decode/Qwen2MoEQ8_0FFNLayersDecode.java b/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/decode/Qwen2MoEQ8_0FFNLayersDecode.java
new file mode 100644
index 00000000..6903e585
--- /dev/null
+++ b/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/decode/Qwen2MoEQ8_0FFNLayersDecode.java
@@ -0,0 +1,105 @@
+package org.beehive.gpullama3.tornadovm.layers.type.q8_0.decode;
+
+import org.beehive.gpullama3.inference.state.Qwen2MoEState;
+import org.beehive.gpullama3.inference.weights.tornado.Qwen2MoETornadoWeights;
+import org.beehive.gpullama3.model.qwen2.Qwen2MoEConfiguration;
+import org.beehive.gpullama3.tornadovm.layers.type.q8_0.Qwen2MoEQ8_0FFNLayers;
+import org.beehive.gpullama3.tornadovm.scheduling.SchedulerType;
+
+import uk.ac.manchester.tornado.api.TaskGraph;
+import uk.ac.manchester.tornado.api.enums.DataTransferMode;
+
+/** Single-token decode layers that continue from the batch-prefill KV cache. */
+public final class Qwen2MoEQ8_0FFNLayersDecode extends Qwen2MoEQ8_0FFNLayers {
+
+ public Qwen2MoEQ8_0FFNLayersDecode(
+ String taskGraph,
+ Qwen2MoEState state,
+ Qwen2MoETornadoWeights weights,
+ Qwen2MoEConfiguration config,
+ SchedulerType schedulerType) {
+ super(taskGraph, state, weights, config, schedulerType);
+ }
+
+ /** Connects decode layer 0 to the decode activation and later layers in order. */
+ @Override
+ protected String predecessorGraphName(int layerIndex) {
+ return layerIndex == 0 ? "decodeActivation" : "layer_" + (layerIndex - 1);
+ }
+
+ /** Reuses the weights already uploaded by the corresponding batch-prefill layer. */
+ @Override
+ protected TaskGraph configureLayerWeights(TaskGraph layer, int layerIndex) {
+ return layer.consumeFromDevice(
+ "batchPrefillLayer_" + layerIndex,
+ weights.rms_att_weightLayered[layerIndex].asFloatArray(),
+ weights.wqLayered[layerIndex].asByteArray(),
+ weights.wkLayered[layerIndex].asByteArray(),
+ weights.wvLayered[layerIndex].asByteArray(),
+ weights.woLayered[layerIndex].asByteArray(),
+ weights.q_biasLayered[layerIndex].asFloatArray(),
+ weights.k_biasLayered[layerIndex].asFloatArray(),
+ weights.v_biasLayered[layerIndex].asFloatArray(),
+ weights.rms_ffn_weightLayered[layerIndex].asFloatArray(),
+ weights.routerGateLayered[layerIndex].asFloatArray(),
+ weights.gateExpertsLayered[layerIndex].asByteArray(),
+ weights.upExpertsLayered[layerIndex].asByteArray(),
+ weights.downExpertsLayered[layerIndex].asByteArray(),
+ weights.sharedGateLayered[layerIndex].asByteArray(),
+ weights.sharedUpLayered[layerIndex].asByteArray(),
+ weights.sharedDownLayered[layerIndex].asByteArray(),
+ weights.sharedGateInputLayered[layerIndex].asFloatArray());
+ }
+
+ /** Reuses the KV cache produced by batch prefill instead of allocating a new cache. */
+ @Override
+ protected TaskGraph configureLayerDataTransfers(TaskGraph layer, int layerIndex) {
+ if (layerIndex == 0) {
+ layer.transferToDevice(
+ DataTransferMode.EVERY_EXECUTION,
+ moeState.positionHolder,
+ moeState.temp,
+ moeState.tempFFN);
+ layer.transferToDevice(
+ DataTransferMode.FIRST_EXECUTION,
+ context,
+ moeState.wrapXb,
+ moeState.wrapXb2,
+ moeState.wrapQ,
+ moeState.wrapK,
+ moeState.wrapV,
+ moeState.wrapAtt,
+ moeState.wrapRouterLogits,
+ moeState.wrapSelectedExperts,
+ moeState.wrapRoutingWeights,
+ moeState.wrapExpertGate,
+ moeState.wrapSharedGate,
+ moeState.wrapSharedOutput);
+ layer.consumeFromDevice(
+ "decodeActivation", moeState.wrapKeyCache, moeState.wrapValueCache);
+ } else {
+ String predecessor = "layer_" + (layerIndex - 1);
+ layer.consumeFromDevice(
+ predecessor,
+ context,
+ moeState.wrapXb,
+ moeState.wrapXb2,
+ moeState.wrapQ,
+ moeState.wrapK,
+ moeState.wrapV,
+ moeState.wrapKeyCache,
+ moeState.wrapValueCache,
+ moeState.wrapAtt,
+ moeState.wrapRouterLogits,
+ moeState.wrapSelectedExperts,
+ moeState.wrapRoutingWeights,
+ moeState.wrapExpertGate,
+ moeState.wrapSharedGate,
+ moeState.wrapSharedOutput,
+ moeState.positionHolder,
+ moeState.temp,
+ moeState.tempFFN);
+ }
+ return layer;
+ }
+}
diff --git a/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/prefill/Qwen2MoEQ8_0LayersBatchPrefill.java b/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/prefill/Qwen2MoEQ8_0LayersBatchPrefill.java
new file mode 100644
index 00000000..1b159f0b
--- /dev/null
+++ b/src/main/java/org/beehive/gpullama3/tornadovm/layers/type/q8_0/prefill/Qwen2MoEQ8_0LayersBatchPrefill.java
@@ -0,0 +1,490 @@
+package org.beehive.gpullama3.tornadovm.layers.type.q8_0.prefill;
+
+import org.beehive.gpullama3.inference.state.Qwen2MoEState;
+import org.beehive.gpullama3.inference.weights.tornado.Qwen2MoETornadoWeights;
+import org.beehive.gpullama3.model.qwen2.Qwen2MoEConfiguration;
+import org.beehive.gpullama3.tornadovm.kernels.Qwen2MoEBatchKernels;
+import org.beehive.gpullama3.tornadovm.kernels.TransformerBatchPrefillKernels;
+import org.beehive.gpullama3.tornadovm.layers.BatchPrefillTransformerLayerTaskGraphs;
+import org.beehive.gpullama3.tornadovm.scheduling.WorkerGridFactory;
+
+import uk.ac.manchester.tornado.api.GridScheduler;
+import uk.ac.manchester.tornado.api.ImmutableTaskGraph;
+import uk.ac.manchester.tornado.api.KernelContext;
+import uk.ac.manchester.tornado.api.TaskGraph;
+import uk.ac.manchester.tornado.api.WorkerGrid;
+import uk.ac.manchester.tornado.api.enums.DataTransferMode;
+
+import java.util.List;
+import java.util.stream.IntStream;
+
+/** Batched-prefill Transformer-layer TaskGraphs for Qwen2-MoE Q8_0. */
+public final class Qwen2MoEQ8_0LayersBatchPrefill
+ implements BatchPrefillTransformerLayerTaskGraphs {
+
+ private static final int LOCAL_WORK_GROUP_SIZE = 32;
+
+ private final Qwen2MoEState state;
+ private final Qwen2MoETornadoWeights weights;
+ private final Qwen2MoEConfiguration config;
+ private final KernelContext context = new KernelContext();
+ private final int batchSize;
+ private final int dim;
+ private final int kvDim;
+ private final int topK;
+ private final int numberOfAssignments;
+ private final List layerTaskGraphs;
+ private String lastLayerTaskGraphID;
+
+ public Qwen2MoEQ8_0LayersBatchPrefill(
+ Qwen2MoEState state,
+ Qwen2MoETornadoWeights weights,
+ Qwen2MoEConfiguration config,
+ int batchSize) {
+ this.state = state;
+ this.weights = weights;
+ this.config = config;
+ this.batchSize = batchSize;
+ this.dim = config.dim();
+ this.kvDim = config.kvDim();
+ this.topK = config.numberOfExpertsUsed();
+ this.numberOfAssignments = batchSize * topK;
+ this.layerTaskGraphs =
+ IntStream.range(0, config.numberOfLayers())
+ .mapToObj(this::createBatchPrefillLayerTaskGraph)
+ .map(TaskGraph::snapshot)
+ .toList();
+ }
+
+ /** Creates the complete batch-prefill TaskGraph for one Transformer layer. */
+ private TaskGraph createBatchPrefillLayerTaskGraph(int layerIndex) {
+ String graphName = "batchPrefillLayer_" + layerIndex;
+ if (layerIndex == config.numberOfLayers() - 1) {
+ lastLayerTaskGraphID = graphName;
+ }
+
+ TaskGraph layer = new TaskGraph(graphName);
+ configureDataTransfers(layer, layerIndex);
+ configureAttention(layer, layerIndex);
+ configureMoE(layer, layerIndex);
+ layer.persistOnDevice(
+ state.wrapXBatch,
+ state.wrapKeyCache,
+ state.wrapValueCache,
+ weights.rms_att_weightLayered[layerIndex].asFloatArray(),
+ weights.wqLayered[layerIndex].asByteArray(),
+ weights.wkLayered[layerIndex].asByteArray(),
+ weights.wvLayered[layerIndex].asByteArray(),
+ weights.woLayered[layerIndex].asByteArray(),
+ weights.q_biasLayered[layerIndex].asFloatArray(),
+ weights.k_biasLayered[layerIndex].asFloatArray(),
+ weights.v_biasLayered[layerIndex].asFloatArray(),
+ weights.rms_ffn_weightLayered[layerIndex].asFloatArray(),
+ weights.routerGateLayered[layerIndex].asFloatArray(),
+ weights.gateExpertsLayered[layerIndex].asByteArray(),
+ weights.upExpertsLayered[layerIndex].asByteArray(),
+ weights.downExpertsLayered[layerIndex].asByteArray(),
+ weights.sharedGateLayered[layerIndex].asByteArray(),
+ weights.sharedUpLayered[layerIndex].asByteArray(),
+ weights.sharedDownLayered[layerIndex].asByteArray(),
+ weights.sharedGateInputLayered[layerIndex].asFloatArray());
+ return layer;
+ }
+
+ /** Declares layer weights and the batch buffers that remain on the GPU. */
+ private void configureDataTransfers(TaskGraph layer, int layerIndex) {
+ if (layerIndex == 0) {
+ layer.transferToDevice(
+ DataTransferMode.EVERY_EXECUTION,
+ state.batchStartPosHolder,
+ state.activeBatchSizeHolder);
+ layer.transferToDevice(
+ DataTransferMode.FIRST_EXECUTION,
+ context,
+ state.attnScaleBatch,
+ state.ffnScaleBatch,
+ state.wrapXbBatch,
+ state.wrapQBatch,
+ state.wrapKBatch,
+ state.wrapVBatch,
+ state.wrapKeyCache,
+ state.wrapValueCache,
+ state.wrapRouterLogitsBatch,
+ state.wrapSelectedExpertsBatch,
+ state.wrapRoutingWeightsBatch,
+ state.wrapGroupedAssignmentIds,
+ state.wrapGroupedPositionByAssignment,
+ state.wrapGroupedExpertHidden,
+ state.wrapGroupedExpertDown,
+ state.wrapSharedHiddenBatch,
+ state.wrapSharedWeightBatch);
+ layer.consumeFromDevice("prefillActivation", state.wrapXBatch);
+ } else {
+ String predecessor = "batchPrefillLayer_" + (layerIndex - 1);
+ layer.consumeFromDevice(
+ predecessor,
+ context,
+ state.wrapXBatch,
+ state.wrapXbBatch,
+ state.wrapQBatch,
+ state.wrapKBatch,
+ state.wrapVBatch,
+ state.wrapKeyCache,
+ state.wrapValueCache,
+ state.batchStartPosHolder,
+ state.activeBatchSizeHolder,
+ state.attnScaleBatch,
+ state.ffnScaleBatch,
+ state.wrapRouterLogitsBatch,
+ state.wrapSelectedExpertsBatch,
+ state.wrapRoutingWeightsBatch,
+ state.wrapGroupedAssignmentIds,
+ state.wrapGroupedPositionByAssignment,
+ state.wrapGroupedExpertHidden,
+ state.wrapGroupedExpertDown,
+ state.wrapSharedHiddenBatch,
+ state.wrapSharedWeightBatch);
+ }
+
+ layer.transferToDevice(
+ DataTransferMode.FIRST_EXECUTION,
+ weights.rms_att_weightLayered[layerIndex].asFloatArray(),
+ weights.wqLayered[layerIndex].asByteArray(),
+ weights.wkLayered[layerIndex].asByteArray(),
+ weights.wvLayered[layerIndex].asByteArray(),
+ weights.q_biasLayered[layerIndex].asFloatArray(),
+ weights.k_biasLayered[layerIndex].asFloatArray(),
+ weights.v_biasLayered[layerIndex].asFloatArray(),
+ weights.woLayered[layerIndex].asByteArray(),
+ weights.rms_ffn_weightLayered[layerIndex].asFloatArray(),
+ weights.routerGateLayered[layerIndex].asFloatArray(),
+ weights.gateExpertsLayered[layerIndex].asByteArray(),
+ weights.upExpertsLayered[layerIndex].asByteArray(),
+ weights.downExpertsLayered[layerIndex].asByteArray(),
+ weights.sharedGateLayered[layerIndex].asByteArray(),
+ weights.sharedUpLayered[layerIndex].asByteArray(),
+ weights.sharedDownLayered[layerIndex].asByteArray(),
+ weights.sharedGateInputLayered[layerIndex].asFloatArray());
+ }
+
+ /** Adds Qwen2 attention tasks for every token in the prefill batch. */
+ private void configureAttention(TaskGraph layer, int layerIndex) {
+ layer.task(
+ "batch_attn_rms",
+ TransformerBatchPrefillKernels::batchedRmsReduceParallel,
+ context,
+ state.wrapXBatch,
+ state.attnScaleBatch,
+ dim,
+ config.rmsNormEps(),
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_attn_rms_apply",
+ TransformerBatchPrefillKernels::batchedRmsApplyFP32,
+ context,
+ state.wrapXbBatch,
+ state.wrapXBatch,
+ weights.rms_att_weightLayered[layerIndex].asFloatArray(),
+ state.attnScaleBatch,
+ dim);
+
+ layer.task(
+ "batch_qkv",
+ TransformerBatchPrefillKernels::batchedFusedQKVMatmulQ8,
+ context,
+ state.wrapXbBatch,
+ state.wrapQBatch,
+ state.wrapKBatch,
+ state.wrapVBatch,
+ weights.wqLayered[layerIndex].asByteArray(),
+ weights.wkLayered[layerIndex].asByteArray(),
+ weights.wvLayered[layerIndex].asByteArray(),
+ dim,
+ kvDim,
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_qkv_bias",
+ Qwen2MoEBatchKernels::batchedQKVBias,
+ context,
+ state.wrapQBatch,
+ state.wrapKBatch,
+ state.wrapVBatch,
+ weights.q_biasLayered[layerIndex].asFloatArray(),
+ weights.k_biasLayered[layerIndex].asFloatArray(),
+ weights.v_biasLayered[layerIndex].asFloatArray(),
+ state.activeBatchSizeHolder,
+ dim,
+ kvDim);
+
+ layer.task(
+ "batch_rope_kv",
+ Qwen2MoEBatchKernels::batchedRopeWithKVCacheQwen2,
+ context,
+ state.batchStartPosHolder,
+ state.activeBatchSizeHolder,
+ state.wrapQBatch,
+ state.wrapKBatch,
+ state.wrapVBatch,
+ state.wrapKeyCache,
+ state.wrapValueCache,
+ kvDim,
+ config.headSize(),
+ layerIndex,
+ config.contextLength(),
+ dim,
+ config.ropeTheta());
+
+ layer.task(
+ "batch_attention",
+ TransformerBatchPrefillKernels::batchedFlashAttention,
+ context,
+ state.batchStartPosHolder,
+ state.wrapQBatch,
+ state.wrapKeyCache,
+ state.wrapValueCache,
+ state.wrapXbBatch,
+ config.numberOfHeads(),
+ config.headSize(),
+ kvDim,
+ config.kvMul(),
+ layerIndex,
+ config.contextLength(),
+ dim);
+
+ layer.task(
+ "batch_attn_out",
+ TransformerBatchPrefillKernels::batchedMatVecWithResidualQ8,
+ context,
+ state.wrapXbBatch,
+ state.wrapXBatch,
+ weights.woLayered[layerIndex].asByteArray(),
+ dim,
+ dim,
+ LOCAL_WORK_GROUP_SIZE);
+ }
+
+ /** Adds batched routing, routed experts, and the shared expert. */
+ private void configureMoE(TaskGraph layer, int layerIndex) {
+ layer.task(
+ "batch_ffn_rms",
+ TransformerBatchPrefillKernels::batchedRmsReduceParallel,
+ context,
+ state.wrapXBatch,
+ state.ffnScaleBatch,
+ dim,
+ config.rmsNormEps(),
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_ffn_rms_apply",
+ TransformerBatchPrefillKernels::batchedRmsApplyFP32,
+ context,
+ state.wrapXbBatch,
+ state.wrapXBatch,
+ weights.rms_ffn_weightLayered[layerIndex].asFloatArray(),
+ state.ffnScaleBatch,
+ dim);
+
+ layer.task(
+ "batch_router",
+ Qwen2MoEBatchKernels::batchedRouterProjection,
+ context,
+ state.wrapXbBatch,
+ state.wrapRouterLogitsBatch,
+ weights.routerGateLayered[layerIndex].asFloatArray(),
+ state.activeBatchSizeHolder,
+ dim,
+ config.numberOfExperts(),
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_topk",
+ Qwen2MoEBatchKernels::batchedSoftmaxAndTopK,
+ context,
+ state.wrapRouterLogitsBatch,
+ state.wrapSelectedExpertsBatch,
+ state.wrapRoutingWeightsBatch,
+ state.activeBatchSizeHolder,
+ config.numberOfExperts(),
+ topK);
+
+ layer.task(
+ "batch_group_assignments",
+ Qwen2MoEBatchKernels::groupAssignmentsByExpert,
+ context,
+ state.wrapSelectedExpertsBatch,
+ state.wrapGroupedAssignmentIds,
+ state.wrapGroupedPositionByAssignment,
+ state.activeBatchSizeHolder,
+ config.numberOfExperts(),
+ topK);
+
+ layer.task(
+ "batch_routed_gate_up",
+ Qwen2MoEBatchKernels::groupedRoutedExpertsGateUpSwiGLUQ8_0,
+ context,
+ state.wrapXbBatch,
+ state.wrapSelectedExpertsBatch,
+ state.wrapGroupedAssignmentIds,
+ state.activeBatchSizeHolder,
+ weights.gateExpertsLayered[layerIndex].asByteArray(),
+ weights.upExpertsLayered[layerIndex].asByteArray(),
+ state.wrapGroupedExpertHidden,
+ dim,
+ config.moeHiddenDim(),
+ config.numberOfExperts(),
+ topK,
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_routed_down",
+ Qwen2MoEBatchKernels::groupedRoutedExpertsDownQ8_0,
+ context,
+ state.wrapGroupedExpertHidden,
+ state.wrapSelectedExpertsBatch,
+ state.wrapGroupedAssignmentIds,
+ state.activeBatchSizeHolder,
+ weights.downExpertsLayered[layerIndex].asByteArray(),
+ state.wrapGroupedExpertDown,
+ dim,
+ config.moeHiddenDim(),
+ config.numberOfExperts(),
+ topK,
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_routed_accumulate",
+ Qwen2MoEBatchKernels::accumulateGroupedRoutedExperts,
+ context,
+ state.wrapGroupedExpertDown,
+ state.wrapGroupedPositionByAssignment,
+ state.wrapRoutingWeightsBatch,
+ state.wrapXBatch,
+ state.activeBatchSizeHolder,
+ dim,
+ topK);
+
+ layer.task(
+ "batch_shared_gate_up",
+ Qwen2MoEBatchKernels::batchedSharedExpertGateUpSwiGLUQ8_0,
+ context,
+ state.wrapXbBatch,
+ state.activeBatchSizeHolder,
+ weights.sharedGateLayered[layerIndex].asByteArray(),
+ weights.sharedUpLayered[layerIndex].asByteArray(),
+ state.wrapSharedHiddenBatch,
+ dim,
+ config.sharedExpertHiddenDim(),
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_shared_weight",
+ Qwen2MoEBatchKernels::batchedSharedExpertGateWeight,
+ context,
+ state.wrapXbBatch,
+ weights.sharedGateInputLayered[layerIndex].asFloatArray(),
+ state.wrapSharedWeightBatch,
+ state.activeBatchSizeHolder,
+ dim,
+ LOCAL_WORK_GROUP_SIZE);
+
+ layer.task(
+ "batch_shared_down",
+ Qwen2MoEBatchKernels::batchedSharedExpertDownAndAccumulateQ8_0,
+ context,
+ state.wrapSharedHiddenBatch,
+ state.wrapSharedWeightBatch,
+ state.activeBatchSizeHolder,
+ weights.sharedDownLayered[layerIndex].asByteArray(),
+ state.wrapXBatch,
+ dim,
+ config.sharedExpertHiddenDim(),
+ LOCAL_WORK_GROUP_SIZE);
+ }
+
+ /** Configures the fixed WorkerGrid used by every task in every layer. */
+ @Override
+ public void updateGridScheduler(GridScheduler scheduler) {
+ WorkerGrid rmsWorker = groupedRowsWorker(batchSize);
+ WorkerGrid batchScalarWorker = WorkerGridFactory.genericWorker(batchSize, 1);
+ WorkerGrid batchElementWorker = WorkerGridFactory.genericWorker(batchSize * dim, 256);
+ WorkerGrid qkvWorker = groupedRowsWorker(batchSize * (dim + 2 * kvDim));
+ int qkvBiasGlobalWork = batchSize * (dim + 2 * kvDim);
+ WorkerGrid qkvBiasWorker =
+ WorkerGridFactory.genericWorker(
+ qkvBiasGlobalWork,
+ validLocalSize(qkvBiasGlobalWork, 256));
+
+ int ropeGlobalWork = batchSize * (dim / 2);
+ int ropeLocalWork = validLocalSize(ropeGlobalWork, 512);
+ WorkerGrid ropeWorker = WorkerGridFactory.genericWorker(ropeGlobalWork, ropeLocalWork);
+
+ int attentionLocalWork = validLocalSize(config.headSize(), 64);
+ WorkerGrid attentionWorker =
+ WorkerGridFactory.genericWorker(
+ batchSize * config.numberOfHeads() * attentionLocalWork,
+ attentionLocalWork);
+
+ WorkerGrid batchDimRowsWorker = groupedRowsWorker(batchSize * dim);
+ WorkerGrid routerWorker = groupedRowsWorker(batchSize * config.numberOfExperts());
+ WorkerGrid groupingWorker = WorkerGridFactory.createSingleWorker();
+ WorkerGrid routedHiddenWorker =
+ groupedRowsWorker(numberOfAssignments * config.moeHiddenDim());
+ WorkerGrid routedDownWorker = groupedRowsWorker(numberOfAssignments * dim);
+ WorkerGrid sharedHiddenWorker =
+ groupedRowsWorker(batchSize * config.sharedExpertHiddenDim());
+ WorkerGrid sharedWeightWorker = groupedRowsWorker(batchSize);
+
+ for (int layer = 0; layer < config.numberOfLayers(); layer++) {
+ String prefix = "batchPrefillLayer_" + layer + ".";
+ scheduler.addWorkerGrid(prefix + "batch_attn_rms", rmsWorker);
+ scheduler.addWorkerGrid(prefix + "batch_attn_rms_apply", batchElementWorker);
+ scheduler.addWorkerGrid(prefix + "batch_qkv", qkvWorker);
+ scheduler.addWorkerGrid(prefix + "batch_qkv_bias", qkvBiasWorker);
+ scheduler.addWorkerGrid(prefix + "batch_rope_kv", ropeWorker);
+ scheduler.addWorkerGrid(prefix + "batch_attention", attentionWorker);
+ scheduler.addWorkerGrid(prefix + "batch_attn_out", batchDimRowsWorker);
+ scheduler.addWorkerGrid(prefix + "batch_ffn_rms", rmsWorker);
+ scheduler.addWorkerGrid(prefix + "batch_ffn_rms_apply", batchElementWorker);
+ scheduler.addWorkerGrid(prefix + "batch_router", routerWorker);
+ scheduler.addWorkerGrid(prefix + "batch_topk", batchScalarWorker);
+ scheduler.addWorkerGrid(prefix + "batch_group_assignments", groupingWorker);
+ scheduler.addWorkerGrid(prefix + "batch_routed_gate_up", routedHiddenWorker);
+ scheduler.addWorkerGrid(prefix + "batch_routed_down", routedDownWorker);
+ scheduler.addWorkerGrid(prefix + "batch_routed_accumulate", batchElementWorker);
+ scheduler.addWorkerGrid(prefix + "batch_shared_gate_up", sharedHiddenWorker);
+ scheduler.addWorkerGrid(prefix + "batch_shared_weight", sharedWeightWorker);
+ scheduler.addWorkerGrid(prefix + "batch_shared_down", batchDimRowsWorker);
+ }
+ }
+
+ /** Creates a 32-thread work-group for every logical output row. */
+ private static WorkerGrid groupedRowsWorker(int rows) {
+ return WorkerGridFactory.genericWorker(rows * LOCAL_WORK_GROUP_SIZE, LOCAL_WORK_GROUP_SIZE);
+ }
+
+ /** Finds a legal local size that divides the requested global dimension. */
+ private static int validLocalSize(int size, int maximum) {
+ int localSize = Math.min(size, maximum);
+ while (localSize > 1 && size % localSize != 0) {
+ localSize--;
+ }
+ return localSize;
+ }
+
+ @Override
+ public List getLayerImmutableTaskGraphs() {
+ return layerTaskGraphs;
+ }
+
+ @Override
+ public String getLastLayerTaskGraphID() {
+ return lastLayerTaskGraphID;
+ }
+
+ public KernelContext getContext() {
+ return context;
+ }
+}
diff --git a/src/main/java/org/beehive/gpullama3/tornadovm/plan/ForwardPlanFactory.java b/src/main/java/org/beehive/gpullama3/tornadovm/plan/ForwardPlanFactory.java
index 6e8ba29d..d96e83e5 100644
--- a/src/main/java/org/beehive/gpullama3/tornadovm/plan/ForwardPlanFactory.java
+++ b/src/main/java/org/beehive/gpullama3/tornadovm/plan/ForwardPlanFactory.java
@@ -179,9 +179,12 @@ private static ForwardPlan createQwen2Q8_0Plan(ExecutionMode mode, Qwen2State st
}
private static ForwardPlan createQwen2MoEQ8_0Plan(ExecutionMode mode, Qwen2MoEState state, Model model) {
- if (mode != ExecutionMode.STANDARD)
- throw new UnsupportedOperationException(mode + " not yet supported for QWEN_2_MOE + Q8_0");
- return new SingleTokenForwardPlan(model, new Qwen2MoEQ8_0PlanComponents(state, model));
+ BatchPrefillDecodeForwardPlanComponents components = new Qwen2MoEQ8_0PlanComponents(state, model);
+ return switch (mode) {
+ case STANDARD -> new SingleTokenForwardPlan(model, components);
+ case PREFILL_DECODE -> throw new UnsupportedOperationException(mode + " not yet supported for QWEN_2_MOE + Q8_0");
+ case BATCH_PREFILL_DECODE -> new BatchPrefillDecodeForwardPlan(model, components, TornadoVMMasterPlan.PREFILL_BATCH_SIZE);
+ };
}
private static ForwardPlan createQwen3FP16Plan(ExecutionMode mode, Qwen3State state, Model model) {
diff --git a/src/main/java/org/beehive/gpullama3/tornadovm/plan/components/q8_0/Qwen2MoEQ8_0PlanComponents.java b/src/main/java/org/beehive/gpullama3/tornadovm/plan/components/q8_0/Qwen2MoEQ8_0PlanComponents.java
index 93a1443f..ea17283c 100644
--- a/src/main/java/org/beehive/gpullama3/tornadovm/plan/components/q8_0/Qwen2MoEQ8_0PlanComponents.java
+++ b/src/main/java/org/beehive/gpullama3/tornadovm/plan/components/q8_0/Qwen2MoEQ8_0PlanComponents.java
@@ -7,15 +7,21 @@
import org.beehive.gpullama3.tornadovm.layers.AbstractLogitsTaskGraph;
import org.beehive.gpullama3.tornadovm.layers.Activation;
import org.beehive.gpullama3.tornadovm.layers.ActivationTaskGraph;
+import org.beehive.gpullama3.tornadovm.layers.BatchPrefillTransformerLayerTaskGraphs;
import org.beehive.gpullama3.tornadovm.layers.TransformerLayerTaskGraphs;
import org.beehive.gpullama3.tornadovm.layers.type.q8_0.LogitsQ8_0Layer;
import org.beehive.gpullama3.tornadovm.layers.type.q8_0.Qwen2MoEQ8_0FFNLayers;
-import org.beehive.gpullama3.tornadovm.plan.components.SingleTokenForwardPlanComponents;
+import org.beehive.gpullama3.tornadovm.layers.type.q8_0.decode.LogitsQ8_0LayerDecode;
+import org.beehive.gpullama3.tornadovm.layers.type.q8_0.decode.Qwen2MoEQ8_0FFNLayersDecode;
+import org.beehive.gpullama3.tornadovm.layers.type.q8_0.prefill.Qwen2MoEQ8_0LayersBatchPrefill;
+import org.beehive.gpullama3.tornadovm.plan.components.BatchPrefillDecodeForwardPlanComponents;
+import org.beehive.gpullama3.tornadovm.plan.components.activation.BatchDecodeActivation;
+import org.beehive.gpullama3.tornadovm.plan.components.activation.BatchPrefillActivation;
import org.beehive.gpullama3.tornadovm.scheduling.SchedulerDetectionService;
import org.beehive.gpullama3.tornadovm.scheduling.SchedulerType;
-/** Assembles the single-token Q8_0 GPU components for Qwen2-MoE. */
-public final class Qwen2MoEQ8_0PlanComponents implements SingleTokenForwardPlanComponents {
+/** Assembles the single-token and batch-prefill Q8_0 GPU components for Qwen2-MoE. */
+public final class Qwen2MoEQ8_0PlanComponents implements BatchPrefillDecodeForwardPlanComponents {
private final Qwen2MoEState state;
private final Qwen2MoETornadoWeights weights;
@@ -34,13 +40,48 @@ public ActivationTaskGraph singleTokenActivation() {
return new Activation("activationUpdate", state, weights, config);
}
+ @Override
+ public ActivationTaskGraph prefillDecodeActivation() {
+ return new Activation("decodeActivation", state, weights, config);
+ }
+
+ @Override
+ public ActivationTaskGraph batchPrefillActivation(int batchSize) {
+ return new BatchPrefillActivation(state, config, batchSize, true);
+ }
+
+ @Override
+ public ActivationTaskGraph batchDecodeActivation(String lastBatchLayerId) {
+ return new BatchDecodeActivation(state, config, lastBatchLayerId, true);
+ }
+
@Override
public TransformerLayerTaskGraphs singleTokenTransformerLayers() {
return new Qwen2MoEQ8_0FFNLayers("qwen2MoEFFN", state, weights, config, schedulerType);
}
+ @Override
+ public TransformerLayerTaskGraphs prefillDecodeTransformerLayers() {
+ return new Qwen2MoEQ8_0FFNLayers("decode", state, weights, config, schedulerType);
+ }
+
+ @Override
+ public TransformerLayerTaskGraphs batchDecodeTransformerLayers() {
+ return new Qwen2MoEQ8_0FFNLayersDecode("decode", state, weights, config, schedulerType);
+ }
+
+ @Override
+ public BatchPrefillTransformerLayerTaskGraphs batchPrefillTransformerLayers(int batchSize) {
+ return new Qwen2MoEQ8_0LayersBatchPrefill(state, weights, config, batchSize);
+ }
+
@Override
public AbstractLogitsTaskGraph singleTokenLogits(String previousGraphId) {
return new LogitsQ8_0Layer("logits", state, weights, config, previousGraphId, schedulerType);
}
+
+ @Override
+ public AbstractLogitsTaskGraph decodeLogits(String previousGraphId) {
+ return new LogitsQ8_0LayerDecode("logits", state, weights, config, previousGraphId, schedulerType);
+ }
}