wandres.dev
GPGPU Y ML · Matrices, convoluciones y tiling

Convoluciones en compute: halo, separabilidad e im2col

El tile con borde de halo, la descomposición separable con sus números, y cuándo conviene convertir una convolución en una multiplicación de matrices.

⏱ 21 min

Una convolución es una GEMM disfrazada, y esa observación tiene dos consecuencias opuestas. La primera es que todo lo que aprendiste sobre reutilización y tiling se aplica igual, con la diferencia de que aquí los datos que se reutilizan son los píxeles del solape entre ventanas vecinas. La segunda es que, literalmente, se puede reescribir una convolución como una multiplicación de matrices y aprovechar un kernel de GEMM ya optimizado. Cuál de las dos vías gana depende del tamaño del núcleo y del número de canales, y la respuesta es distinta para un desenfoque y para una capa de una red neuronal.

🎯 Al terminar esta lección sabrás
  • Cargar un tile con halo en memoria compartida y contar los accesos que ahorra.
  • Descomponer un núcleo separable y calcular la reducción de operaciones.
  • Implementar un desenfoque gaussiano separable en dos dispatches.
  • Decidir entre convolución directa e im2col con GEMM.

El tile con halo

Para calcular un tile de T por T píxeles de salida con un núcleo de radio R, hacen falta (T + 2R) por (T + 2R) píxeles de entrada: el tile más un borde de anchura R a cada lado. Ese borde se llama halo.

La cuenta de reutilización es directa. Sin memoria compartida, cada uno de los píxeles de salida lee (2R+1)² píxeles de entrada. Con el tile en memoria compartida, se leen (T+2R)² de memoria global y el resto de accesos van a memoria compartida.

T R lecturas sin tile lecturas con tile factor
16 1 2.304 324 7,1
16 2 6.400 400 16,0
16 4 20.736 576 36,0
32 4 82.944 1.600 51,8

El factor crece con el radio y con el lado del tile. Y hay un efecto secundario que se nota: la fracción de halo, que es lo que se carga de más respecto al tile útil. Con T = 16 y R = 4, se cargan 576 píxeles para producir 256: un 125 por ciento de sobrecarga. Con T = 32 y el mismo radio, 1600 para producir 1024: un 56 por ciento. Los tiles grandes amortizan mejor el halo, y ese es el argumento a favor de subir el tamaño del tile hasta donde permita el límite de invocaciones por workgroup.

const T: u32 = 16u;
const R: u32 = 2u;
const LADO: u32 = T + 2u * R;      // 20

@group(0) @binding(0) var origen:  texture_2d<f32>;
@group(0) @binding(1) var destino: texture_storage_2d<rgba8unorm, write>;
@group(0) @binding(2) var<uniform> nucleo: array<vec4f, 7>;   // 25 pesos

var<workgroup> tile: array<vec4f, LADO * LADO>;

@compute @workgroup_size(T, T)
fn convolucion(@builtin(workgroup_id)        wid: vec3u,
               @builtin(local_invocation_id) lid: vec3u,
               @builtin(local_invocation_index) li: u32) {
  let dims = vec2i(textureDimensions(origen));
  let esquina = vec2i(wid.xy) * i32(T) - vec2i(i32(R));

  // 256 invocaciones cargan 400 pixeles: bucle con zancada.
  var k = li;
  loop {
    if (k >= LADO * LADO) { break; }
    let p = esquina + vec2i(i32(k % LADO), i32(k / LADO));
    let q = clamp(p, vec2i(0), dims - vec2i(1));   // repetir el borde
    tile[k] = textureLoad(origen, q, 0);
    k += T * T;
  }
  workgroupBarrier();

  var suma = vec4f(0.0);
  for (var dy: u32 = 0u; dy <= 2u * R; dy = dy + 1u) {
    for (var dx: u32 = 0u; dx <= 2u * R; dx = dx + 1u) {
      let idx = dy * (2u * R + 1u) + dx;
      let peso = nucleo[idx / 4u][idx % 4u];
      suma += tile[(lid.y + dy) * LADO + (lid.x + dx)] * peso;
    }
  }

  let salida = vec2i(wid.xy) * i32(T) + vec2i(lid.xy);
  if (salida.x < dims.x && salida.y < dims.y) {
    textureStore(destino, salida, suma);
  }
}

El bucle de carga con zancada T*T es lo que reparte 400 cargas entre 256 invocaciones sin divergencia estructural y manteniendo la contigüidad. Y el clamp a los bordes de la imagen evita el escalón oscuro que produciría leer ceros fuera.

Fíjate en que el guardia de salida está al final y no al principio: hay una barrera de por medio, así que ninguna invocación puede salir antes.

Separabilidad

Un núcleo bidimensional es separable si se puede escribir como el producto exterior de dos núcleos unidimensionales. El gaussiano lo es, la media lo es, el de Sobel lo es. Un núcleo arbitrario no.

Cuando lo es, la convolución de (2R+1)² multiplicaciones se sustituye por dos pasadas de (2R+1) cada una:

R directa separable factor
1 9 6 1,5
2 25 10 2,5
4 81 18 4,5
8 289 34 8,5

El factor es aproximadamente (2R+1)/2, y para radios grandes es enorme. El precio es una pasada extra sobre la imagen, con su lectura y escritura completas, y una textura intermedia.

// Pasada horizontal: solo el eje X necesita halo.
const T: u32 = 64u;
const R: u32 = 4u;
const LADO: u32 = T + 2u * R;

@group(0) @binding(0) var origen:     texture_2d<f32>;
@group(0) @binding(1) var intermedia: texture_storage_2d<rgba16float, write>;
@group(0) @binding(2) var<uniform> pesos: array<vec4f, 3>;   // 9 pesos 1D

fn peso1D(d: u32) -> f32 { return pesos[d / 4u][d % 4u]; }

var<workgroup> fila: array<vec4f, LADO>;

@compute @workgroup_size(T)
fn horizontal(@builtin(workgroup_id) wid: vec3u,
              @builtin(local_invocation_index) li: u32) {
  let dims = vec2i(textureDimensions(origen));
  let y = i32(wid.y);
  let x0 = i32(wid.x * T) - i32(R);

  var k = li;
  loop {
    if (k >= LADO) { break; }
    let x = clamp(x0 + i32(k), 0, dims.x - 1);
    fila[k] = textureLoad(origen, vec2i(x, y), 0);
    k += T;
  }
  workgroupBarrier();

  var suma = vec4f(0.0);
  for (var d: u32 = 0u; d <= 2u * R; d = d + 1u) {
    suma += fila[li + d] * peso1D(d);
  }

  let x = i32(wid.x * T) + i32(li);
  if (x < dims.x) { textureStore(intermedia, vec2i(x, y), suma); }
}

La pasada vertical es la misma con los ejes cambiados, leyendo de intermedia. Y ahí hay una trampa de rendimiento: la pasada vertical, escrita de forma simétrica, lee columnas, y una columna es el peor patrón de acceso posible. Las implementaciones serias hacen la pasada vertical sobre un tile bidimensional en vez de sobre una tira, o transponen la imagen intermedia y aplican dos veces el mismo kernel horizontal. La transposición cuesta una pasada más y suele salir a cuenta en radios grandes.

im2col y el kernel de GEMM

En una red neuronal, una capa convolucional aplica Cout filtros de Cin canales y k por k píxeles sobre un mapa de activaciones. Escrito como convolución directa, cada invocación acumula Cin·k² productos, y con Cin = 256 y k = 3 son 2304 por píxel de salida.

La transformación im2col reordena las activaciones de forma que cada ventana de k por k por Cin se convierte en una columna de una matriz. Entonces la convolución entera es una multiplicación de matrices: la matriz de filtros, de Cout por Cin·k², por la matriz de columnas, de Cin·k² por el número de píxeles de salida.

La ventaja es contundente: puedes usar el kernel de GEMM con tiling en memoria compartida y en registros que ya has escrito, que llega al 60 o 70 por ciento del pico, en vez de escribir un kernel de convolución específico que rara vez pasa del 30. La desventaja es que la matriz de columnas ocupa veces más memoria que las activaciones originales —nueve veces con k = 3— y hay que construirla, lo que cuesta una pasada de escritura.

El criterio práctico:

Convolución directa cuando el núcleo es pequeño y los canales pocos, cuando la memoria es escasa, o cuando la convolución es por canal —las depthwise, donde cada canal se convoluciona por separado y no hay ninguna reducción entre canales que convertir en un producto de matrices—. Los modelos móviles modernos están llenos de convoluciones por canal, y para ellas im2col no aporta nada.

im2col más GEMM cuando hay muchos canales, que es el caso de las capas intermedias de cualquier red convolucional seria. A partir de 64 canales de entrada la GEMM gana casi siempre.

Hay una tercera vía que evita la memoria extra: im2col implícito, donde el kernel de GEMM lee directamente de las activaciones calculando el índice de la ventana en lugar de leer de una matriz materializada. Es lo que hacen las bibliotecas serias, y es bastante más código porque la aritmética de índices se mete dentro del bucle de carga del tile.

La separabilidad se puede verificar, y muchos nucleos que parecen no serlo lo son aproximadamente

Un núcleo bidimensional es exactamente separable si su matriz tiene rango uno, y eso se comprueba con una descomposición en valores singulares en tres líneas de cualquier lenguaje: si solo hay un valor singular no nulo, es separable, y los dos vectores unidimensionales son el primer vector singular por la izquierda y por la derecha escalados por la raíz del valor singular. Lo interesante es lo que ocurre cuando hay varios valores singulares no nulos pero uno domina: el núcleo no es separable, pero se puede aproximar por la suma de dos o tres pares separables, con un error que se conoce exactamente porque son los valores singulares descartados. Un núcleo de 15 por 15 arbitrario cuesta 225 multiplicaciones por píxel; aproximado con los tres primeros términos de su descomposición cuesta 3 por 30, o sea 90, y el error relativo es la energía de los valores singulares restantes. Para un desenfoque de lente, un bokeh con forma de hexágono o un núcleo de dispersión subsuperficial, dos o tres términos suelen bastar para que la diferencia sea invisible. Es una técnica que se usa poco fuera del cine y que en la web tiene todo el sentido, porque el cuello de botella es exactamente el número de accesos por píxel. El procedimiento completo: calculas la descomposición en JavaScript al cargar, decides cuántos términos según un umbral de energía, generas los pares de vectores, y ejecutas dos pasadas separables por término acumulando en el destino. Y hay un caso especialmente rentable: los núcleos de convolución aprendidos por una red, que suelen tener rango efectivo muy bajo porque el entrenamiento los regulariza, y donde esta descomposición es una forma de compresión del modelo que además acelera la inferencia.