wandres.dev
COMPUTE III · Reducciones y prefix sum

Scan de varios bloques y compactación de arrays

Las tres pasadas que escanean un array de cualquier tamaño, la recursión sobre las sumas de bloque, y el filtro paralelo que se construye encima.

⏱ 22 min

Un workgroup escanea 512 elementos. Un array de un millón necesita 1954 bloques, y cada bloque produce un scan correcto en su interior y completamente equivocado respecto al array entero, porque le falta sumar todo lo que había antes. Corregir ese desfase es un problema de scan sobre las sumas de bloque, o sea el mismo problema una escala más arriba. De esa observación sale la estructura de tres pasadas que resuelve el scan de cualquier tamaño, y con ella el filtro paralelo, que es la operación que de verdad usa todo el mundo.

🎯 Al terminar esta lección sabrás
  • Encadenar las tres pasadas del scan de bloques con las dependencias correctas.
  • Calcular cuántos niveles de recursión hacen falta según el tamaño de entrada.
  • Implementar una compactación de array completa sobre el scan.
  • Situar el scan de una sola pasada y las garantías que le faltan en WebGPU.

Las tres pasadas

La descomposición es directa. Sea B el tamaño de bloque, en nuestro caso 512.

Pasada uno. Cada workgroup hace el scan exclusivo de su bloque en memoria compartida y escribe dos cosas: el scan local en el array de salida, y la suma total de su bloque en un array auxiliar de numBloques posiciones. Es exactamente el kernel de Blelloch, que ya escribe las dos cosas.

Pasada dos. Se escanea el array de sumas de bloque, en exclusivo. La posición k de ese scan contiene la suma de todos los bloques anteriores al k, que es justo el desfase que le falta al bloque k.

Pasada tres. Cada elemento del bloque k suma su desfase. Es un kernel elemento a elemento trivial.

flowchart TB
E[Array de entrada de n elementos] --> P1[Pasada 1 scan por bloque]
P1 --> S1[Scan local incompleto en cada bloque]
P1 --> SB[Array de sumas de bloque]
SB --> P2[Pasada 2 scan del array de sumas]
P2 --> DE[Desfase acumulado por bloque]
S1 --> P3[Pasada 3 sumar el desfase]
DE --> P3
P3 --> R[Scan global correcto]
style E fill:#89b4fa,color:#11111b
style P1 fill:#fab387,color:#11111b
style P2 fill:#fab387,color:#11111b
style P3 fill:#fab387,color:#11111b
style S1 fill:#94e2d5,color:#11111b
style SB fill:#94e2d5,color:#11111b
style DE fill:#94e2d5,color:#11111b
style R fill:#a6e3a1,color:#11111b

El kernel de la tercera pasada:

const N: u32 = 512u;   // elementos por bloque, igual que en la pasada 1

struct Params { conteo: u32 };
@group(0) @binding(0) var<uniform> params: Params;
@group(0) @binding(1) var<storage, read_write> datos:    array<f32>;
@group(0) @binding(2) var<storage, read>       desfases: array<f32>;

@compute @workgroup_size(256)
fn sumarDesfase(@builtin(global_invocation_id) gid: vec3u) {
  let i = gid.x;
  if (i >= params.conteo) { return; }
  datos[i] = datos[i] + desfases[i / N];
}

Aquí sí vale el return temprano: no hay ninguna barrera en el kernel.

Y la orquestación, con los tres dispatches en el mismo pass:

const N = 512;                                   // elementos por bloque
const bloques = Math.ceil(n / N);

const ST = GPUBufferUsage.STORAGE;
const CD = GPUBufferUsage.COPY_DST;
const CS = GPUBufferUsage.COPY_SRC;

const bufDatos  = device.createBuffer({ size: n * 4,       usage: ST | CD | CS });
const bufSalida = device.createBuffer({ size: n * 4,       usage: ST | CS });
const bufSumas  = device.createBuffer({ size: bloques * 4, usage: ST });
const bufDesf   = device.createBuffer({ size: bloques * 4, usage: ST });

const enc = device.createCommandEncoder();
const pass = enc.beginComputePass();

// 1. scan por bloque -> bufSalida, y totales -> bufSumas
pass.setPipeline(pipeScanBloque);
pass.setBindGroup(0, bgPasada1);
pass.dispatchWorkgroups(bloques);

// 2. scan de bufSumas -> bufDesf   (un solo workgroup si bloques <= N)
pass.setPipeline(pipeScanBloque);
pass.setBindGroup(0, bgPasada2);
pass.dispatchWorkgroups(Math.ceil(bloques / N));

// 3. bufSalida[i] += bufDesf[i / N]
pass.setPipeline(pipeSumarDesfase);
pass.setBindGroup(0, bgPasada3);
pass.dispatchWorkgroups(Math.ceil(n / 256));

pass.end();
device.queue.submit([enc.finish()]);

Cuántos niveles de recursión

La pasada dos es un scan, así que sufre el mismo problema si el array de sumas no cabe en un bloque. Con B = 512:

elementos bloques niveles
hasta 512 1 1
hasta 262.144 hasta 512 2
hasta 134.217.728 hasta 262.144 3

Tres niveles cubren 134 millones de elementos, que con f32 son 512 MiB y superan el límite por defecto de tamaño de buffer. En la práctica, nunca hacen falta más de tres, y para casi cualquier caso real bastan dos.

Eso permite escribir la recursión como un bucle acotado en lugar de como una función recursiva, y encolar todos los dispatches de antemano:

function encolarScan(pass, buffers, n) {
  const niveles = [];
  let conteo = n;
  while (conteo > 1) {
    niveles.push(conteo);
    conteo = Math.ceil(conteo / N);
  }
  // Bajada: scan por bloque en cada nivel.
  for (let k = 0; k < niveles.length; k++) {
    pass.setPipeline(pipeScanBloque);
    pass.setBindGroup(0, buffers.bajada[k]);
    pass.dispatchWorkgroups(Math.ceil(niveles[k] / N));
  }
  // Subida: sumar desfases en orden inverso, saltando el nivel mas alto.
  for (let k = niveles.length - 2; k >= 0; k--) {
    pass.setPipeline(pipeSumarDesfase);
    pass.setBindGroup(0, buffers.subida[k]);
    pass.dispatchWorkgroups(Math.ceil(niveles[k] / 256));
  }
}

El coste total en operaciones es 2n para la primera bajada más 2n/B para la segunda, más n para la subida: del orden de 3n, o sea lineal. El coste en tráfico de memoria es mayor y suele ser lo que manda: se lee y escribe el array entero en la pasada uno y otra vez en la tres, unas cuatro veces n en total, contra las dos veces n que sería el mínimo teórico.

El filtro paralelo

El uso más frecuente del scan no es sumar números, es asignar posiciones. Compactar un array —quedarse solo con los elementos que cumplen una condición, en el mismo orden— son tres kernels y el scan en medio.

// Kernel 1: marcar. bandera[i] = 1 si el elemento sobrevive.
@group(0) @binding(0) var<uniform> params: Params;
@group(0) @binding(1) var<storage, read>       fuente:  array<Particula>;
@group(0) @binding(2) var<storage, read_write> bandera: array<u32>;

@compute @workgroup_size(256)
fn marcar(@builtin(global_invocation_id) gid: vec3u) {
  let i = gid.x;
  if (i >= params.conteo) { return; }
  bandera[i] = select(0u, 1u, fuente[i].vida > 0.0);
}
// Kernel 3: dispersar. destino[posicion[i]] = fuente[i] si sobrevive.
@group(0) @binding(1) var<storage, read>       fuente:   array<Particula>;
@group(0) @binding(2) var<storage, read>       bandera:  array<u32>;
@group(0) @binding(3) var<storage, read>       posicion: array<u32>;   // scan exclusivo
@group(0) @binding(4) var<storage, read_write> destino:  array<Particula>;
@group(0) @binding(5) var<storage, read_write> total:    array<u32>;

@compute @workgroup_size(256)
fn dispersar(@builtin(global_invocation_id) gid: vec3u) {
  let i = gid.x;
  if (i >= params.conteo) { return; }
  if (bandera[i] == 1u) {
    destino[posicion[i]] = fuente[i];
  }
  // La ultima invocacion deja escrito cuantos sobrevivieron.
  if (i == params.conteo - 1u) {
    total[0] = posicion[i] + bandera[i];
  }
}

El total que deja la última invocación es el número de supervivientes, y se puede usar directamente como argumento de un dispatchWorkgroupsIndirect o de un drawIndirect sin que el dato pase nunca por la CPU. Es el patrón que hace posible el reciclado de partículas.

La propiedad que distingue esta compactación de la versión con atomicAdd es que conserva el orden. Un contador atómico también reparte posiciones únicas, pero en el orden en que las invocaciones llegan a la atómica, que depende de la planificación y cambia entre ejecuciones. Si el array estaba ordenado por profundidad para dibujar transparencias, la versión atómica lo desordena y la del scan no.

El scan de una sola pasada

Las tres pasadas leen y escriben el array completo dos veces, y hay una técnica que lo hace en una: el decoupled look-back. Cada bloque publica primero su suma local con una marca de “agregado disponible”, luego mira hacia atrás sumando los agregados de los bloques anteriores hasta encontrar uno que ya tenga el prefijo total, y entonces publica su propio prefijo total. El tráfico baja a lo mínimo teórico y en hardware nativo es la implementación estándar.

En WebGPU tiene un problema conceptual serio: el mecanismo de mirar hacia atrás requiere que un bloque espere a que otro publique, y eso es precisamente lo que la frontera entre workgroups no garantiza. Si el bloque al que esperas no ha sido planificado porque tú ocupas su sitio, el bucle no termina.

Hay implementaciones reales en producción que lo usan, apoyándose en que los planificadores actuales despachan los workgroups en orden creciente, lo cual hace que un bloque nunca espere a otro posterior. Es una suposición razonable sobre el hardware existente y no una garantía de la especificación, y se ha discutido activamente en el grupo de trabajo. La postura sensata para código que va a ejecutarse en máquinas desconocidas es usar las tres pasadas, que son correctas por construcción, y reservar el look-back para cuando midas que el scan es tu cuello de botella y tengas un plan de repliegue.

El desfase de la pasada tres se puede fusionar con la pasada uno, y ahi hay un 25 por ciento gratis

La estructura de tres pasadas tiene una ineficiencia que se ve en cuanto se cuenta el tráfico: la pasada uno escribe el array entero y la pasada tres lo vuelve a leer entero para sumarle una constante por bloque y escribirlo otra vez. Son dos lecturas y dos escrituras de n elementos cuando el algoritmo solo necesita, conceptualmente, una de cada. La fusión que casi nadie aplica consiste en no escribir el resultado en la pasada uno: la pasada uno solo calcula las sumas de bloque, que son n/B valores, y las escribe. Después de la pasada dos, un único kernel vuelve a leer la entrada, recalcula el scan local del bloque en memoria compartida —el mismo trabajo de antes— y escribe directamente el resultado ya desfasado. Se ha duplicado el cálculo del scan local y se ha eliminado una escritura y una lectura completas del array. En un scan de flotantes, donde el operador es una suma y el kernel está limitado por memoria de principio a fin, ese cambio mide entre un 20 y un 30 por ciento de mejora, y la razón es que en una GPU moderna una suma en registros cuesta del orden de un ciclo y un acceso a VRAM cuesta cientos. Es el mismo principio que hace ganar al tiling en la multiplicación de matrices, y es la lección transversal de todo el cómputo en GPU: recalcular sale más barato que releer, casi siempre, y por márgenes que en la CPU serían impensables.