wandres.dev
COMPUTE III · Reducciones y prefix sum

La reducción en árbol dentro de un workgroup

El primer ladrillo del cómputo paralelo: coste, profundidad, el direccionamiento que evita divergencia y la corrección para tamaños que no son potencia de dos.

⏱ 22 min

Reducir es convertir un array en un solo valor con un operador asociativo: sumar, multiplicar, buscar el máximo, contar los que cumplen algo. En una CPU es un bucle de tres líneas y coste lineal. En una GPU es la estructura de datos fundamental de la que salen el scan, la ordenación por radix, la construcción de rejillas espaciales y la normalización de un softmax. Merece que se le dedique el rigor de un capítulo de algoritmia: cuánto trabajo cuesta, cuántos pasos ocupa, y por qué la versión evidente es dos veces más lenta que la correcta.

🎯 Al terminar esta lección sabrás
  • Derivar el trabajo total y la profundidad de una reducción en árbol.
  • Escribir una reducción de workgroup correcta con direccionamiento secuencial.
  • Garantizar el resultado para tamaños de entrada que no son potencia de dos.
  • Aplicar el teorema de Brent para elegir cuántos elementos procesa cada invocación.

La estructura y su coste

La idea es sustituir el bucle secuencial de n pasos por un árbol de profundidad logarítmica. En cada paso, la mitad de las invocaciones activas combinan su valor con el de otra y el número de valores vivos se divide por dos.

flowchart TB
A0[d0] --> B0[d0 mas d4]
A4[d4] --> B0
A1[d1] --> B1[d1 mas d5]
A5[d5] --> B1
A2[d2] --> B2[d2 mas d6]
A6[d6] --> B2
A3[d3] --> B3[d3 mas d7]
A7[d7] --> B3
B0 --> C0[suma de 0 2 4 6]
B2 --> C0
B1 --> C1[suma de 1 3 5 7]
B3 --> C1
C0 --> D0[suma total en la posicion 0]
C1 --> D0
style A0 fill:#89b4fa,color:#11111b
style A1 fill:#89b4fa,color:#11111b
style A2 fill:#89b4fa,color:#11111b
style A3 fill:#89b4fa,color:#11111b
style A4 fill:#89b4fa,color:#11111b
style A5 fill:#89b4fa,color:#11111b
style A6 fill:#89b4fa,color:#11111b
style A7 fill:#89b4fa,color:#11111b
style B0 fill:#94e2d5,color:#11111b
style B1 fill:#94e2d5,color:#11111b
style B2 fill:#94e2d5,color:#11111b
style B3 fill:#94e2d5,color:#11111b
style C0 fill:#fab387,color:#11111b
style C1 fill:#fab387,color:#11111b
style D0 fill:#a6e3a1,color:#11111b

El trabajo total es el número de operaciones binarias ejecutadas. En el primer paso se hacen n/2, en el segundo n/4, y así hasta una. La suma es n/2 + n/4 + ... + 1 = n - 1. Exactamente las mismas operaciones que el bucle secuencial, ni una más: la reducción en árbol es óptima en trabajo.

La profundidad es el número de pasos con dependencia entre ellos, y es log2(n). Para 256 elementos son 8 pasos; para un millón serían 20. Cada paso necesita una barrera, porque el valor que una invocación lee lo acaba de escribir otra.

De ahí sale el atractivo: el tiempo pasa de lineal a logarítmico sin gastar más operaciones. Pero hay una letra pequeña. En el paso k solo hay n/2^k invocaciones activas: en el último paso trabaja una y las 255 restantes están paradas dentro de la barrera. El número medio de invocaciones ocupadas a lo largo del árbol es aproximadamente n / log2(n), o sea 32 de 256. La eficiencia paralela es baja, y eso condiciona el diseño de la última sección.

La versión correcta

const TAM: u32 = 256u;

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> parciales: array<f32>;

var<workgroup> buf: array<f32, TAM>;

@compute @workgroup_size(TAM)
fn main(@builtin(global_invocation_id) gid: vec3u,
        @builtin(local_invocation_index) li:  u32,
        @builtin(workgroup_id)           wid: vec3u) {

  // 1. Carga con relleno neutro: nadie sale por return, todos llegan
  //    a todas las barreras.
  var v = 0.0;                       // 0 es el neutro de la suma
  if (gid.x < params.conteo) { v = entrada[gid.x]; }
  buf[li] = v;
  workgroupBarrier();

  // 2. Arbol con direccionamiento secuencial.
  for (var s: u32 = TAM / 2u; s > 0u; s >>= 1u) {
    if (li < s) { buf[li] = buf[li] + buf[li + s]; }
    workgroupBarrier();
  }

  // 3. Una escritura por workgroup.
  if (li == 0u) { parciales[wid.x] = buf[0]; }
}

Tres decisiones merecen justificación.

El relleno neutro en vez del return. Es lo que hace que el kernel sea correcto para cualquier conteo, sea o no potencia de dos. El árbol siempre recorre TAM posiciones, y TAM sí es potencia de dos porque es el tamaño del workgroup, que tú eliges. Las posiciones sobrantes contienen el elemento neutro del operador y no alteran el resultado. Un return temprano rompería la uniformidad de las barreras y dejaría el resultado indefinido, como vimos al colocar el guardia.

El neutro cambia con el operador y hay que acertarlo:

operador neutro
suma 0.0
producto 1.0
máximo -3.4028235e38
mínimo 3.4028235e38
y lógico true
o lógico false
y bit a bit 0xffffffffu
o bit a bit 0u

Para máximo y mínimo se usa el flotante finito de mayor magnitud en vez del infinito, porque WGSL no admite literales infinitos. Funciona siempre que ningún dato de entrada lo supere, cosa que no puede ocurrir con f32.

El direccionamiento secuencial. La versión que sale sola cuando uno dibuja el árbol es la intercalada: en el paso s, la invocación li actúa si li % (2*s) == 0 y combina buf[li] con buf[li + s]. Produce el resultado correcto y es notablemente peor.

// PEOR: direccionamiento intercalado.
for (var s: u32 = 1u; s < TAM; s <<= 1u) {
  if (li % (2u * s) == 0u) { buf[li] += buf[li + s]; }
  workgroupBarrier();
}

Con direccionamiento intercalado, las invocaciones activas en el primer paso son las pares: dentro de un warp de 32, la mitad trabaja y la mitad no, así que el warp entero se ejecuta y solo aprovecha la mitad de sus carriles. Al llegar al paso 5, la única invocación activa de cada warp obliga a mantener vivo el warp completo. Con direccionamiento secuencial, en cambio, las invocaciones activas son siempre las primeras s, que son contiguas: en cuanto s baja de 256 a 128, hay warps enteros sin nada que hacer y el planificador los retira. La divergencia dentro del warp desaparece. Además, buf[li] y buf[li + s] con li contiguo reparte los accesos entre todos los bancos de memoria compartida, mientras que el paso 2*s del intercalado los concentra.

La diferencia medida entre las dos versiones ronda el factor dos, y no cuesta nada elegir la buena.

El operador dentro de una función. Si vas a escribir varias reducciones, mete la combinación en una función y el neutro en una constante. El compilador la incorpora y no cuesta nada:

const NEUTRO: f32 = 0.0;
fn comb(a: f32, b: f32) -> f32 { return a + b; }

El teorema de Brent y los elementos por invocación

Volvamos al problema de la eficiencia. Con 256 invocaciones y 256 elementos, la reducción hace 255 operaciones repartidas en 8 pasos, pero ocupa 256 invocaciones durante los 8 pasos: 2048 ranuras de ejecución para 255 operaciones útiles, un 12 por ciento de aprovechamiento.

El teorema de Brent dice que un algoritmo con trabajo W y profundidad D se puede ejecutar en p procesadores en tiempo W/p + D. Traducido: en vez de asignar una invocación por elemento, asigna a cada invocación un tramo del array que reduce secuencialmente, y solo después haz el árbol. El tiempo pasa a ser n/p + log2(p), y con p mucho menor que n el segundo término desaparece.

@compute @workgroup_size(TAM)
fn main(@builtin(global_invocation_id) gid: vec3u,
        @builtin(local_invocation_index) li:  u32,
        @builtin(workgroup_id)           wid: vec3u,
        @builtin(num_workgroups)         nwg: vec3u) {

  // Fase secuencial: cada invocacion suma varios elementos con zancada
  // igual al total de invocaciones, para no perder coalescencia.
  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) { parciales[wid.x] = buf[0]; }
}

Ahora el número de workgroups no lo dicta el tamaño del array sino la máquina: se lanza un número que llene el chip —del orden de cuatro a ocho workgroups por unidad de cómputo, o sea unos cientos— y cada invocación absorbe tantos elementos como haga falta. Con 16 millones de elementos y 512 workgroups de 256, cada invocación suma 122 elementos en la fase secuencial y luego participa en 8 pasos de árbol. La fase secuencial ocupa el cien por cien de los carriles y domina el tiempo; el árbol es un remate barato.

Hay una ventaja añadida que no es de rendimiento: el número de parciales que hay que consolidar en el segundo dispatch baja de decenas de miles a unos cientos.

El bucle con zancada, ojo, contiene un break que depende de los datos. Es perfectamente legal porque no hay ninguna barrera dentro del bucle. La uniformidad solo se exige donde hay barreras.

La reduccion en arbol es mas precisa que el bucle secuencial, y eso cambia como se valida

Existe la creencia extendida de que un cálculo en GPU es menos preciso que el mismo cálculo en CPU. Para una reducción de flotantes es exactamente al revés, y con una diferencia que se puede cuantificar. El error de redondeo acumulado de una suma secuencial de n términos crece como O(n·eps) en el peor caso, porque cada suma parcial es cada vez mayor que el término que se le añade y los bits bajos del sumando pequeño se pierden. En una reducción en árbol, cada suma combina dos valores de magnitud parecida y la profundidad es logarítmica, así que el error crece como O(log(n)·eps). Con un millón de sumandos positivos de magnitud similar, la diferencia entre lineal y logarítmica son cinco órdenes de magnitud de error. En la práctica esto significa que si comparas tu reducción de GPU contra un array.reduce((a,b)=>a+b, 0) de JavaScript y no coinciden, el que probablemente está mal es el de JavaScript, y compararlos con una tolerancia estrecha te llevará a perseguir un bug que no existe. La forma correcta de validar es calcular la referencia con suma compensada de Kahan o con la suma en doble precisión, y comparar contra eso; entonces se ve que la reducción en árbol se le acerca mucho más que el bucle ingenuo. El mismo argumento se vuelve en tu contra en el caso opuesto: si tu kernel usa una atómica en coma flotante para acumular, el orden de llegada es arbitrario y pierdes tanto la reproducibilidad como la garantía de magnitudes parecidas. Por eso, en cualquier reducción donde la precisión importe, el árbol con dos dispatches gana al truco de la atómica en las dos cosas que importan.