java-topology/defects/vllm/test/VllmLoraConvertMappingTest.java
russell@unturf.com 80b26e82b6 ml-inference scan: vllm-0001 LoRA convert_mapping list.index O(B*L); ollama+langchain CLEAN
vllm-0001: punica_wrapper/utils.py convert_mapping() calls
lora_index_to_id.index(x) per token in batch — O(B*L) where
B=batch_size, L=loaded_loras. Code has "TODO index can be slow"
comment. Fix: pre-build dict for O(1) lookup. 7x at B=2000/L=64.

Ollama: all slices.Contains on bounded slices (1-8 items).
LangChain: orchestration code, all membership bounded by k param.
2026-03-30 16:41:47 -04:00

135 lines
5.3 KiB
Java

/**
* CWE-407 unit test for vllm-0001: LoRA convert_mapping list.index() O(B*L)
*
* Defect: vllm/lora/punica_wrapper/utils.py convert_mapping()
* - lora_index_to_id.index(x) called inside loop over index_mapping_indices
* - O(B * L) where B = batch_size (token count), L = loaded LoRA count
* - Code even has "# TODO index can be slow. optimize" comment
*
* Fix: Pre-build dict {lora_id: index} for O(1) lookup → O(B + L)
*/
import java.util.*;
public class VllmLoraConvertMappingTest {
// --- Defective: O(B * L) list.index() per token ---
static int[] convertMappingDefective(int[] indexMappingIndices, List<Integer> loraIndexToId) {
int[] loraIndices = new int[indexMappingIndices.length];
for (int i = 0; i < indexMappingIndices.length; i++) {
if (indexMappingIndices[i] > 0) {
// Linear scan: O(L) per token
loraIndices[i] = loraIndexToId.indexOf(indexMappingIndices[i]);
} else {
loraIndices[i] = -1;
}
}
return loraIndices;
}
// --- Fixed: O(B + L) with pre-built HashMap ---
static int[] convertMappingFixed(int[] indexMappingIndices, List<Integer> loraIndexToId) {
// Pre-build reverse lookup: lora_id -> position
Map<Integer, Integer> idToIndex = new HashMap<>();
for (int i = 0; i < loraIndexToId.size(); i++) {
Integer v = loraIndexToId.get(i);
if (v != null && v > 0) {
idToIndex.put(v, i);
}
}
int[] loraIndices = new int[indexMappingIndices.length];
for (int i = 0; i < indexMappingIndices.length; i++) {
if (indexMappingIndices[i] > 0) {
loraIndices[i] = idToIndex.getOrDefault(indexMappingIndices[i], -1);
} else {
loraIndices[i] = -1;
}
}
return loraIndices;
}
public static void main(String[] args) {
System.out.println("=== vllm-0001: LoRA convert_mapping list.index() O(B*L) ===");
// Test correctness first
testCorrectness();
// Benchmark at various scales
int[] batchSizes = {100, 500, 2000, 8000};
int[] loraCounts = {8, 16, 32, 64};
System.out.printf("\n%-12s %-10s %-14s %-14s %-10s%n",
"BatchSize", "LoRAs", "Defective(us)", "Fixed(us)", "Ratio");
System.out.println("-".repeat(62));
for (int batchSize : batchSizes) {
for (int loraCount : loraCounts) {
// Build lora_index_to_id: [None, lora1, lora2, ..., loraN]
List<Integer> loraIndexToId = new ArrayList<>();
loraIndexToId.add(null); // slot 0 = no LoRA
for (int j = 1; j <= loraCount; j++) {
loraIndexToId.add(j * 100); // LoRA IDs: 100, 200, ...
}
// Build index_mapping_indices: tokens assigned to various LoRAs
Random rng = new Random(42);
int[] indices = new int[batchSize];
for (int i = 0; i < batchSize; i++) {
// ~20% no-LoRA, rest distributed across LoRAs
if (rng.nextDouble() < 0.2) {
indices[i] = 0;
} else {
indices[i] = (rng.nextInt(loraCount) + 1) * 100;
}
}
// Warmup
for (int w = 0; w < 50; w++) {
convertMappingDefective(indices, loraIndexToId);
convertMappingFixed(indices, loraIndexToId);
}
// Benchmark defective
int iters = 2000;
long startDef = System.nanoTime();
for (int it = 0; it < iters; it++) {
convertMappingDefective(indices, loraIndexToId);
}
long defectiveNs = System.nanoTime() - startDef;
// Benchmark fixed
long startFix = System.nanoTime();
for (int it = 0; it < iters; it++) {
convertMappingFixed(indices, loraIndexToId);
}
long fixedNs = System.nanoTime() - startFix;
double defUs = defectiveNs / 1000.0 / iters;
double fixUs = fixedNs / 1000.0 / iters;
double ratio = defUs / fixUs;
System.out.printf("%-12d %-10d %-14.1f %-14.1f %-10.1fx%n",
batchSize, loraCount, defUs, fixUs, ratio);
}
}
System.out.println("\nAll tests PASSED.");
}
static void testCorrectness() {
List<Integer> loraIndexToId = new ArrayList<>(Arrays.asList(null, 100, 200, 300));
int[] indices = {0, 100, 200, 300, 0, 100};
int[] expected = {-1, 1, 2, 3, -1, 1};
int[] defResult = convertMappingDefective(indices, loraIndexToId);
int[] fixResult = convertMappingFixed(indices, loraIndexToId);
for (int i = 0; i < expected.length; i++) {
assert defResult[i] == expected[i] :
"Defective mismatch at " + i + ": " + defResult[i] + " != " + expected[i];
assert fixResult[i] == expected[i] :
"Fixed mismatch at " + i + ": " + fixResult[i] + " != " + expected[i];
}
System.out.println("Correctness: PASS");
}
}