Reducción entre workgroups: el segundo dispatch
Cómo consolidar los parciales de miles de workgroups en un solo valor con exactamente dos dispatches, reutilizando el mismo kernel.
Un workgroup sabe reducir sus 256 elementos y escribir un parcial. Con dos mil workgroups quedan dos mil parciales y ningún mecanismo que los junte, porque los workgroups no se pueden sincronizar entre sí. La consolidación tiene que ocurrir en otro dispatch, y la pregunta interesante es cuántos hacen falta: la respuesta, contra lo que sugiere la idea de recursión, es siempre dos.
- Encadenar dos dispatches para reducir un array de cualquier tamaño.
- Escribir un kernel genérico que sirva para las dos pasadas sin duplicar código.
- Justificar por qué dos dispatches bastan y cuándo conviene una tercera pasada.
- Decidir entre dejar el resultado en la GPU y traérselo a la CPU.
Por qué dos y no log n
La formulación de libro de una reducción global es recursiva: reduce por bloques, obtén n/256 parciales, vuelve a reducir por bloques, obtén n/65536, y así hasta uno. Con dieciséis millones de elementos salen tres pasadas.
Esa recursión sobra en cuanto la primera pasada usa el bucle de zancada de la lección anterior. Con la zancada, el número de workgroups deja de depender del tamaño del array: se elige para llenar el chip y punto, del orden de 256 o 512. Y un array de 512 parciales lo reduce un único workgroup en la segunda pasada, también con zancada, porque 512 dividido entre 256 invocaciones son dos elementos por invocación.
Dos dispatches, para cualquier tamaño de entrada. Y el segundo con un solo workgroup, que es lo más barato que se puede lanzar.
// El numero de grupos de la primera pasada no depende de n.
// Se elige para llenar el chip: unos pocos por unidad de computo.
const TAM = 256;
const GRUPOS = 512; // 512 parciales, siempre
Un kernel para las dos pasadas
Las dos pasadas hacen lo mismo: leer un array de f32, reducirlo con zancada, escribir un parcial por workgroup. La única diferencia es qué buffer leen, cuál escriben y cuántos elementos hay. Todo eso cabe en un bind group distinto y un uniform.
const TAM: u32 = 256u;
const NEUTRO: f32 = 0.0;
fn comb(a: f32, b: f32) -> f32 { return a + b; }
struct Params { conteo: u32 };
@group(0) @binding(0) var<uniform> params: Params;
@group(0) @binding(1) var<storage, read> entrada: array<f32>;
@group(0) @binding(2) var<storage, read_write> salida: array<f32>;
var<workgroup> buf: array<f32, TAM>;
@compute @workgroup_size(TAM)
fn reducir(@builtin(global_invocation_id) gid: vec3u,
@builtin(local_invocation_index) li: u32,
@builtin(workgroup_id) wid: vec3u,
@builtin(num_workgroups) nwg: vec3u) {
let zancada = nwg.x * TAM;
var acc = NEUTRO;
var i = gid.x;
loop {
if (i >= params.conteo) { break; }
acc = comb(acc, entrada[i]);
i += zancada;
}
buf[li] = acc;
workgroupBarrier();
for (var s: u32 = TAM / 2u; s > 0u; s >>= 1u) {
if (li < s) { buf[li] = comb(buf[li], buf[li + s]); }
workgroupBarrier();
}
if (li == 0u) { salida[wid.x] = buf[0]; }
}
Fíjate en que el kernel no sabe en qué pasada está. En la primera, entrada es el array de datos y salida el de parciales; en la segunda se intercambian los papeles y salida es un buffer de un solo elemento.
El montaje en JavaScript, con los dos bind groups creados una vez y reutilizados en todos los frames:
const TAM = 256, GRUPOS = 512;
const bufDatos = device.createBuffer({
size: n * 4,
usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST,
});
const bufParciales = device.createBuffer({
size: GRUPOS * 4,
usage: GPUBufferUsage.STORAGE,
});
const bufTotal = device.createBuffer({
size: 4,
usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC,
});
// Dos uniforms distintos: cada pasada tiene su conteo.
const uniA = device.createBuffer({ size: 16, usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST });
const uniB = device.createBuffer({ size: 16, usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST });
device.queue.writeBuffer(uniA, 0, new Uint32Array([n]));
device.queue.writeBuffer(uniB, 0, new Uint32Array([GRUPOS]));
const layout = pipeline.getBindGroupLayout(0);
const bgA = device.createBindGroup({ layout, entries: [
{ binding: 0, resource: { buffer: uniA } },
{ binding: 1, resource: { buffer: bufDatos } },
{ binding: 2, resource: { buffer: bufParciales } },
]});
const bgB = device.createBindGroup({ layout, entries: [
{ binding: 0, resource: { buffer: uniB } },
{ binding: 1, resource: { buffer: bufParciales } },
{ binding: 2, resource: { buffer: bufTotal } },
]});
const enc = device.createCommandEncoder();
const pass = enc.beginComputePass();
pass.setPipeline(pipeline);
pass.setBindGroup(0, bgA);
pass.dispatchWorkgroups(GRUPOS); // n -> 512 parciales
pass.setBindGroup(0, bgB);
pass.dispatchWorkgroups(1); // 512 -> 1
pass.end();
device.queue.submit([enc.finish()]);
El uniform tiene 16 bytes para un solo u32 porque el tamaño mínimo de un binding uniform está alineado a 16 bytes; los 12 restantes son relleno. Es una de las reglas de alineación que muerden en cuanto uno intenta ahorrar.
El segundo dispatch ve lo que escribió el primero sin ninguna barrera explícita: es la garantía entre dispatches consecutivos del mismo pass que ya vimos al hablar de la frontera entre workgroups.
Cuándo hace falta una tercera
Solo en un caso: cuando el número de parciales es tan grande que la fase secuencial del segundo dispatch domina el tiempo. Con GRUPOS fijado en 512 eso no ocurre nunca. Ocurre si insistes en lanzar un workgroup por cada 256 elementos, porque entonces dieciséis millones de elementos producen 65536 parciales y un único workgroup tendría que sumar 256 elementos por invocación en serie, en un dispatch que no aprovecha más que una unidad de cómputo de las treinta que tiene el chip.
O sea que la tercera pasada es el síntoma de haber elegido mal el número de workgroups de la primera. Si te encuentras necesitándola, la corrección es fijar GRUPOS y añadir el bucle de zancada, no añadir un dispatch.
Elegir GRUPOS con criterio se puede hacer sin adivinar demasiado. Un valor entre 256 y 1024 funciona bien en todo el rango de hardware que soporta WebGPU, porque incluso una GPU integrada modesta tiene unas pocas decenas de unidades de cómputo y una de gama alta unas ciento y pico. La API no expone el número de unidades de cómputo —deliberadamente, por huella digital— así que no hay forma de calcularlo, y un valor fijo generoso es la respuesta pragmática.
Dejarlo en la GPU o traerlo
Si el total lo necesita otro shader —normalizar un vector, escalar un histograma, calcular la media para restarla— no lo traigas. Deja bufTotal como storage y bíndalo en el dispatch siguiente. El coste es cero y evitas la latencia de sincronización.
// Tercer dispatch: usar el total sin que salga de la GPU.
@group(0) @binding(0) var<storage, read> total: array<f32>;
@group(0) @binding(1) var<storage, read_write> datos: array<f32>;
@compute @workgroup_size(256)
fn normalizar(@builtin(global_invocation_id) gid: vec3u) {
if (gid.x >= arrayLength(&datos)) { return; }
let t = total[0];
if (t != 0.0) { datos[gid.x] = datos[gid.x] / t; }
}
Si de verdad lo necesitas en JavaScript —para mostrarlo, para decidir algo— hay que copiar a un buffer con MAP_READ y esperar. El coste no es el de los cuatro bytes: es que mapAsync no resuelve hasta que la GPU ha terminado todo el trabajo encolado, lo que en la práctica introduce al menos un frame de latencia y a menudo dos.
enc.copyBufferToBuffer(bufTotal, 0, bufLectura, 0, 4);
device.queue.submit([enc.finish()]);
await bufLectura.mapAsync(GPUMapMode.READ);
const total = new Float32Array(bufLectura.getMappedRange())[0];
bufLectura.unmap();
El patrón que evita el parón es el de doble buffer con lectura diferida: se copia el total de cada frame a un buffer de una tanda de tres, y se lee el de hace tres frames, que ya está listo. El dato llega tarde pero no bloquea, y para un contador que se muestra en pantalla nadie nota la diferencia.
Hay una técnica conocida para hacer la reducción en un solo dispatch: cada workgroup escribe su parcial y después incrementa un contador atómico global; el último workgroup en llegar —el que recibe como valor antiguo numGrupos - 1— sabe que todos los demás han terminado, y se encarga él solo de reducir el array de parciales. Es ingeniosa, elimina un dispatch, y aparece en muchos artículos. También es frágil por dos motivos que conviene tener presentes antes de copiarla. El primero es de corrección: para que el último workgroup pueda leer los parciales de los demás hace falta que esas escrituras sean visibles, y la única garantía que da WGSL en ese punto es que la atómica es indivisible, no que las escrituras normales previas ya estén publicadas. Hay que ordenar con storageBarrier() antes de la atómica y aun así se está operando en el filo de lo que el modelo de memoria promete. El segundo es de rendimiento: el ahorro real es el sobrecoste de un dispatch, unos pocos microsegundos, mientras que el coste es que un único workgroup hace la segunda fase con el resto del chip parado, exactamente igual que en la versión de dos dispatches, pero además con toda la maquinaria del contador. En una reducción de 16 millones de elementos, la primera fase se lleva más del 95 por ciento del tiempo y el dispatch que ahorras es ruido. La conclusión práctica: dos dispatches, kernel único, y la complejidad que te ahorras la inviertes en algo que se note. La técnica del último workgroup sí tiene su sitio, pero es en el scan de un solo paso, donde la comunicación entre bloques no es un remate sino el corazón del algoritmo.