Bitonic sort en WebGPU: dos kernels y 78 dispatches
El kernel global paso a paso, el kernel de bloque que hace nueve pasos sin salir de memoria compartida, y los offsets dinámicos que evitan un bind group por paso.
La red bitónica de un millón de elementos tiene 210 pasos. Implementada de la forma directa, eso son 210 dispatches encadenados, cada uno leyendo y escribiendo cuatro megabytes de VRAM, y en total unos ocho milisegundos de puro tráfico de memoria en una GPU de escritorio. La versión que se usa de verdad hace 78 dispatches y mueve la cuarta parte de datos, y la diferencia está entera en una observación: cuando la distancia entre los elementos que se comparan es menor que el bloque que cabe en memoria compartida, todos los pasos que quedan se pueden hacer sin salir del workgroup.
- Escribir el kernel de un paso global de la red bitónica.
- Escribir el kernel de bloque que ejecuta varios pasos en memoria compartida.
- Orquestar los dos kernels con offsets dinámicos en un único buffer uniform.
- Rellenar hasta la potencia de dos y ordenar pares de clave y carga.
El kernel global
Un paso de la red es un comparador por cada par de posiciones a distancia j. La invocación toma su índice como el elemento bajo del par y trabaja solo si su compañero está por delante.
struct Paso { j: u32, k: u32, n: u32 };
@group(0) @binding(0) var<uniform> p: Paso;
@group(0) @binding(1) var<storage, read_write> claves: array<u32>;
@compute @workgroup_size(64)
fn pasoGlobal(@builtin(global_invocation_id) gid: vec3u) {
let i = gid.x;
if (i >= p.n) { return; }
let l = i ^ p.j;
if (l <= i) { return; } // cada par lo procesa el indice menor
let creciente = (i & p.k) == 0u;
let a = claves[i];
let b = claves[l];
if ((a > b) == creciente) {
claves[i] = b;
claves[l] = a;
}
}
No hay barreras, así que los return tempranos son legales. La mitad de las invocaciones sale en la segunda comprobación: es un desperdicio del 50 por ciento que se podría evitar lanzando n/2 invocaciones y reconstruyendo el índice, pero como el kernel está completamente limitado por memoria y las invocaciones que salen no leen nada, en la práctica apenas se nota.
Lo que sí importa es el patrón de acceso. Con j grande, claves[i] y claves[l] están muy separados, así que cada warp hace dos lecturas coalescentes en dos zonas distintas de memoria. Es aceptable. Con j pequeño, i y l son casi adyacentes y el acceso es eficiente, pero entonces el kernel global no debería estar ejecutándose: ese caso lo cubre el kernel de bloque.
El kernel de bloque
Un workgroup de 256 invocaciones carga 512 claves en memoria compartida. Mientras la distancia j sea menor que 512, el compañero de cualquier elemento del bloque está dentro del mismo bloque, así que todos esos pasos se pueden ejecutar con barreras internas y una sola lectura y escritura de VRAM.
La aritmética cambia respecto al kernel global. Aquí conviene lanzar exactamente un comparador por invocación, sin desperdiciar la mitad, y para eso hay que mapear el índice de invocación t al elemento bajo del comparador:
i = (t / j) · 2j + (t mod j) l = i + j
Con j = 1 da los pares (0,1), (2,3), (4,5). Con j = 2 da (0,2), (1,3), (4,6), (5,7). Con j = 4 da (0,4), (1,5), (2,6), (3,7), (8,12). Es la enumeración correcta de los comparadores, sin repeticiones y sin huecos, y como j es potencia de dos el compilador convierte la división y el módulo en desplazamientos.
Hay dos variantes del kernel. La primera ordena el bloque entero de cero, cubriendo todas las etapas hasta k = 512:
const WG: u32 = 256u;
const BLOQUE: u32 = 2u * WG; // 512 elementos por workgroup
@group(0) @binding(1) var<storage, read_write> claves: array<u32>;
var<workgroup> s: array<u32, BLOQUE>;
@compute @workgroup_size(WG)
fn ordenarBloque(@builtin(workgroup_id) wid: vec3u,
@builtin(local_invocation_index) t: u32) {
let base = wid.x * BLOQUE;
s[t] = claves[base + t];
s[t + WG] = claves[base + t + WG];
workgroupBarrier();
for (var k: u32 = 2u; k <= BLOQUE; k <<= 1u) {
for (var j: u32 = k >> 1u; j > 0u; j >>= 1u) {
let i = (t / j) * (2u * j) + (t % j);
let l = i + j;
let creciente = ((base + i) & k) == 0u;
let a = s[i];
let b = s[l];
if ((a > b) == creciente) { s[i] = b; s[l] = a; }
workgroupBarrier();
}
}
claves[base + t] = s[t];
claves[base + t + WG] = s[t + WG];
}
Los 45 pasos que corresponden a las etapas de 2 a 512 caben en este único dispatch. Los dos bucles tienen límites constantes, así que la barrera está en control de flujo uniforme; y como cada invocación toca un par de posiciones disjunto del de las demás dentro de un mismo paso, basta una barrera por paso.
El sentido del comparador se calcula con el índice global base + i, no con el local. Es imprescindible: la alternancia creciente-decreciente que construye las bitónicas del nivel siguiente depende de la posición en el array completo.
La segunda variante recibe una etapa k desde el uniform y ejecuta solo los pasos con j menor que el bloque:
struct Paso { j: u32, k: u32, n: u32 };
@group(0) @binding(0) var<uniform> p: Paso;
@compute @workgroup_size(WG)
fn fusionarBloque(@builtin(workgroup_id) wid: vec3u,
@builtin(local_invocation_index) t: u32) {
let base = wid.x * BLOQUE;
s[t] = claves[base + t];
s[t + WG] = claves[base + t + WG];
workgroupBarrier();
for (var j: u32 = BLOQUE >> 1u; j > 0u; j >>= 1u) {
let i = (t / j) * (2u * j) + (t % j);
let l = i + j;
let creciente = ((base + i) & p.k) == 0u;
let a = s[i];
let b = s[l];
if ((a > b) == creciente) { s[i] = b; s[l] = a; }
workgroupBarrier();
}
claves[base + t] = s[t];
claves[base + t + WG] = s[t + WG];
}
Nueve pasos por dispatch en vez de uno.
La orquestación y los offsets dinámicos
Cada paso global necesita valores distintos de j y k en el uniform, y no se puede llamar a writeBuffer en medio de un pass. La solución limpia son los offsets dinámicos: un único buffer uniform con todos los pares precalculados y un setBindGroup que apunta a la entrada que toca.
El desplazamiento debe ser múltiplo de minUniformBufferOffsetAlignment, que por defecto vale 256 bytes. Así que cada entrada ocupa 256 bytes aunque solo use 12.
const WG = 256, BLOQUE = 512, PASO_BYTES = 256;
// 1. Precalcular la lista de pasos globales.
const pasos = [];
for (let k = BLOQUE * 2; k <= n; k <<= 1) {
for (let j = k >> 1; j >= BLOQUE; j >>= 1) pasos.push({ j, k, local: false });
pasos.push({ j: 0, k, local: true }); // fusion de bloque
}
// 2. Un solo buffer uniform con todas las entradas.
const uni = device.createBuffer({
size: pasos.length * PASO_BYTES,
usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST,
});
const datos = new Uint32Array(pasos.length * PASO_BYTES / 4);
pasos.forEach((p, idx) => {
const o = idx * PASO_BYTES / 4;
datos[o] = p.j; datos[o + 1] = p.k; datos[o + 2] = n;
});
device.queue.writeBuffer(uni, 0, datos);
// 3. Layout con offset dinamico y UN solo bind group.
const layout = device.createBindGroupLayout({ entries: [
{ binding: 0, visibility: GPUShaderStage.COMPUTE,
buffer: { type: 'uniform', hasDynamicOffset: true, minBindingSize: 12 } },
{ binding: 1, visibility: GPUShaderStage.COMPUTE,
buffer: { type: 'storage' } },
]});
const bg = device.createBindGroup({ layout, entries: [
{ binding: 0, resource: { buffer: uni, size: 12 } },
{ binding: 1, resource: { buffer: bufClaves } },
]});
// 4. Todo el sort en un solo pass.
const pass = enc.beginComputePass();
pass.setBindGroup(0, bg, [0]);
pass.setPipeline(pipeOrdenarBloque);
pass.dispatchWorkgroups(n / BLOQUE);
pasos.forEach((p, idx) => {
pass.setBindGroup(0, bg, [idx * PASO_BYTES]);
if (p.local) {
pass.setPipeline(pipeFusionarBloque);
pass.dispatchWorkgroups(n / BLOQUE);
} else {
pass.setPipeline(pipePasoGlobal);
pass.dispatchWorkgroups(n / 64);
}
});
pass.end();
Para n de un millón y bloques de 512, la lista tiene 77 entradas y con el dispatch inicial salen 78 dispatches frente a los 210 de la versión ingenua. El ahorro real es mayor que ese cociente, porque los 132 pasos que desaparecen son precisamente los de j pequeño, que en la versión global harían una lectura y una escritura completas del array cada uno.
Relleno y pares clave-carga
La red exige potencia de dos. El buffer se dimensiona a la siguiente potencia y la cola se rellena una sola vez con el centinela máximo, que para orden creciente es 0xffffffff:
const nPot = 1 << Math.ceil(Math.log2(n));
const relleno = new Uint32Array(nPot - n).fill(0xffffffff);
device.queue.writeBuffer(bufClaves, n * 4, relleno);
Ese relleno se escribe una vez si el número de elementos es fijo. Si cambia cada frame, es un kernel de una línea antes del sort.
Para ordenar pares de clave y carga hay dos caminos. El directo es tener dos buffers y aplicar el mismo intercambio a los dos:
@group(0) @binding(1) var<storage, read_write> claves: array<u32>;
@group(0) @binding(2) var<storage, read_write> cargas: array<u32>;
// ... dentro del comparador:
if ((a > b) == creciente) {
claves[i] = b; claves[l] = a;
let ca = cargas[i]; let cb = cargas[l];
cargas[i] = cb; cargas[l] = ca;
}
El otro, más rápido cuando la clave tiene pocos bits significativos, es empaquetar clave y carga en un solo u32: la clave en los bits altos y el índice en los bajos. Con 20 bits de índice quedan 12 para la clave, suficiente para una profundidad cuantizada. Se ordena la mitad de datos, se rompen los empates de forma determinista, y el desempaquetado es una máscara.
fn empaquetar(clave12: u32, indice20: u32) -> u32 {
return (clave12 << 20u) | (indice20 & 0xfffffu);
}
fn indiceDe(v: u32) -> u32 { return v & 0xfffffu; }
Hay una tentación natural al implementar esto: abrir un compute pass por paso, o incluso un command buffer por paso, porque así se puede llamar a writeBuffer entre medias y no hace falta la maquinaria de los offsets dinámicos. Funciona y produce el resultado correcto. También multiplica el coste por un factor que en un móvil llega a diez. La razón es que cada beginComputePass y cada submit obligan al controlador a vaciar y reconfigurar el estado, y un submit además cruza la frontera entre el proceso de la página y el proceso de GPU del navegador, con su serialización de comandos. Setenta y ocho submits en un frame de dieciséis milisegundos no dejan tiempo para nada más. La versión correcta encola los 78 dispatches en un único compute pass dentro de un único command buffer, y ahí el coste por dispatch se reduce a la sincronización interna de la GPU, que es lo mínimo posible. La restricción que eso impone —no poder cambiar un uniform entre dispatches— es exactamente el problema que los offsets dinámicos resuelven, y esa es la razón de que existan, aunque en los tutoriales aparezcan siempre como una optimización de matrices de modelo. Hay una alternativa a los offsets dinámicos que también es válida: crear un bind group por paso durante la inicialización y guardarlos en un array. Cuesta más memoria de objetos de la API y una llamada de creación por paso, pero se hace una sola vez y evita la aritmética de alineación. Lo que no es una alternativa es reconfigurar el uniform a mitad del pass, porque la API sencillamente no lo permite.