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); + } }