Tiling en memoria compartida: el factor 16
El kernel de GEMM por bloques, la cuenta exacta de lecturas que ahorra un tile de 16 por 16, y el tiling en registros que multiplica otra vez.
El tiling es la razón por la que existe la memoria compartida. La idea cabe en una frase —cargar un bloque una vez y usarlo muchas— y su efecto sobre la multiplicación de matrices es exactamente cuantificable: un tile de 16 por 16 divide las lecturas de memoria global por dieciséis, y esa es toda la mejora, ni un porcentaje más ni menos. Es el ejemplo más limpio que conozco de una optimización cuya magnitud se puede demostrar en dos líneas de aritmética antes de escribir el código.
- Derivar el factor exacto de reducción de lecturas de un tile de lado
T. - Implementar el kernel de GEMM por bloques con las barreras y los guardias correctos.
- Comprobar el consumo de memoria compartida frente al límite.
- Añadir tiling en registros y calcular la nueva intensidad aritmética.
La cuenta, antes del código
Un workgroup calcula un bloque de T por T elementos de C. Para hacerlo necesita, en total, T filas completas de A y T columnas completas de B, es decir 2·T·K valores distintos.
Sin memoria compartida, cada una de las T² invocaciones lee sus 2K valores por su cuenta: 2·T²·K lecturas globales.
Con memoria compartida, el bloque de C se calcula en K/T fases. En cada fase se carga un tile de A de T por T y uno de B de T por T, o sea 2T² valores, y con ellos cada invocación hace T productos. El total de lecturas globales es (K/T)·2T² = 2·T·K.
sin tiling: 2 · T² · K lecturas
con tiling: 2 · T · K lecturas
factor: T
El factor de mejora es exactamente el lado del tile. Con T = 16, dieciséis veces menos lecturas de memoria global. Con T = 32, treinta y dos.
Concretamente, para un bloque de 16 por 16 y K = 1024:
| lecturas globales por bloque de 16x16 | |
|---|---|
| sin tiling | 2 · 256 · 1024 = 524.288 |
| con tiling | 2 · 16 · 1024 = 32.768 |
Y la intensidad aritmética sube de 0,25 a 0,25 · 16 = 4 FLOP por byte. Con los 450 GB/s de la máquina del ejemplo, eso son 1800 GFLOP/s previstos, contra los 112 de la versión ingenua.
El kernel
const T: u32 = 16u;
struct Dims { M: u32, N: u32, K: u32, _r: u32 };
@group(0) @binding(0) var<uniform> d: Dims;
@group(0) @binding(1) var<storage, read> A: array<f32>;
@group(0) @binding(2) var<storage, read> B: array<f32>;
@group(0) @binding(3) var<storage, read_write> C: array<f32>;
var<workgroup> tA: array<array<f32, T>, T>;
var<workgroup> tB: array<array<f32, T>, T>;
@compute @workgroup_size(T, T)
fn gemm(@builtin(workgroup_id) wid: vec3u,
@builtin(local_invocation_id) lid: vec3u) {
let filaLocal = lid.y;
let colLocal = lid.x;
let fila = wid.y * T + filaLocal;
let col = wid.x * T + colLocal;
var acc = 0.0;
let fases = (d.K + T - 1u) / T;
for (var f: u32 = 0u; f < fases; f = f + 1u) {
let kA = f * T + colLocal; // columna de A que carga esta invocacion
let kB = f * T + filaLocal; // fila de B que carga esta invocacion
// Relleno con cero fuera del dominio: el producto no se altera y
// nadie sale por return, asi que las barreras siguen siendo uniformes.
var vA = 0.0;
if (fila < d.M && kA < d.K) { vA = A[fila * d.K + kA]; }
var vB = 0.0;
if (kB < d.K && col < d.N) { vB = B[kB * d.N + col]; }
tA[filaLocal][colLocal] = vA;
tB[filaLocal][colLocal] = vB;
workgroupBarrier();
for (var t: u32 = 0u; t < T; t = t + 1u) {
acc = acc + tA[filaLocal][t] * tB[t][colLocal];
}
workgroupBarrier(); // antes de sobrescribir en la fase siguiente
}
if (fila < d.M && col < d.N) {
C[fila * d.N + col] = acc;
}
}
Los puntos que hay que respetar, y que son los mismos de todo el bloque de cómputo:
Dos barreras por fase. La primera separa la carga del cálculo. La segunda separa el cálculo de la carga siguiente, y es la que se olvida: sin ella, una invocación rápida sobrescribe tA mientras otra todavía está en su bucle interno de la fase anterior.
Cero como relleno, nunca un return. El bucle de fases contiene barreras, así que su recorrido tiene que ser idéntico para todas las invocaciones. Cargar cero fuera del dominio hace el kernel correcto para dimensiones arbitrarias, sin exigir que M, N y K sean múltiplos de 16. El cero es el neutro del producto acumulado.
El patrón de carga es coalescente. A[fila * K + kA] con colLocal variando recorre elementos contiguos de una fila de A. B[kB * N + col] con colLocal variando recorre elementos contiguos de una fila de B. Los dos accesos globales son óptimos, que es más de lo que se podía decir de la versión ingenua.
El consumo de memoria compartida es 2 · 16 · 16 · 4 = 2048 bytes, un octavo del límite de 16384. Con T = 32 serían 8192 bytes, todavía dentro; con T = 64 serían 32768 y no cabría, que es el techo estructural del tamaño de tile en WebGPU con los límites por defecto.
Y una restricción que a menudo se olvida: T · T es el número de invocaciones del workgroup, y maxComputeInvocationsPerWorkgroup vale 256. Así que T = 16 es el mayor tile cuadrado posible con un hilo por elemento. Para llegar a tiles mayores hay que hacer que cada invocación calcule varios elementos, que es exactamente lo que viene ahora.
Tiling en registros
Con T = 16 la intensidad es 4 y el punto crítico de la máquina está en 22. Todavía queda un factor de cinco de margen, y para conseguirlo hay que reutilizar en el nivel siguiente de la jerarquía: los registros.
La idea es que cada invocación calcule un pequeño bloque de R por R elementos de C en vez de uno. Los valores que lee de memoria compartida sirven entonces para R productos en vez de uno, y el tráfico entre memoria compartida y registros se divide por R.
const T: u32 = 16u; // invocaciones por lado
const R: u32 = 4u; // elementos de C por invocacion y lado
const TB: u32 = T * R; // 64: lado del bloque de C que calcula el workgroup
var<workgroup> tA: array<array<f32, TB>, TB>; // no cabe: ver mas abajo
Ahí aparece el problema: un tile de 64 por 64 flotantes son 16384 bytes por array, y hacen falta dos. La solución estándar es hacer los tiles rectangulares: el tile de A es de TB por TK y el de B de TK por TB, con TK mucho menor que TB. Con TB = 64 y TK = 8, cada tile ocupa 64·8·4 = 2048 bytes y los dos suman 4096.
const TB: u32 = 64u; // lado del bloque de C
const TK: u32 = 8u; // profundidad de la fase
const R: u32 = 4u; // 16x16 invocaciones x 4x4 elementos = 64x64
var<workgroup> tA: array<f32, TB * TK>;
var<workgroup> tB: array<f32, TK * TB>;
@compute @workgroup_size(16, 16)
fn gemmRegistros(@builtin(workgroup_id) wid: vec3u,
@builtin(local_invocation_id) lid: vec3u) {
var acc: array<array<f32, R>, R>; // 16 acumuladores en registros
for (var i = 0u; i < R; i++) {
for (var j = 0u; j < R; j++) { acc[i][j] = 0.0; }
}
for (var f: u32 = 0u; f < (d.K + TK - 1u) / TK; f = f + 1u) {
// ... cargar tA y tB cooperativamente, con relleno cero ...
workgroupBarrier();
for (var k: u32 = 0u; k < TK; k = k + 1u) {
// R valores de A y R de B, leidos una vez de memoria compartida
var a: array<f32, R>;
var b: array<f32, R>;
for (var i = 0u; i < R; i++) { a[i] = tA[(lid.y * R + i) * TK + k]; }
for (var j = 0u; j < R; j++) { b[j] = tB[k * TB + lid.x * R + j]; }
// R*R = 16 multiplicaciones y sumas con 2R = 8 lecturas
for (var i = 0u; i < R; i++) {
for (var j = 0u; j < R; j++) { acc[i][j] += a[i] * b[j]; }
}
}
workgroupBarrier();
}
// ... escribir los R*R elementos de C con sus guardias ...
}
La cuenta del bucle interno es lo importante: 8 lecturas de memoria compartida producen 16 multiplicaciones y sumas, o sea 32 operaciones. Sin el tiling en registros serían 32 lecturas para las mismas 32 operaciones. Un factor de cuatro más, que sumado al 16 del tiling en memoria compartida da 64.
La intensidad respecto a memoria global sube a 0,25 · 64 = 16 FLOP por byte, ya cerca del punto crítico de 22, y el rendimiento previsto pasa de 1800 a algo del orden de 7000 GFLOP/s. Ese es el territorio donde están las implementaciones serias de GEMM en WebGPU.
Lo que hay que vigilar es la presión de registros: R = 4 significa 16 acumuladores más 8 valores temporales, y a partir de ahí el compilador empieza a volcar registros a memoria y todo el beneficio se evapora. R = 4 es el punto dulce en la mayoría del hardware; R = 8 casi siempre es peor.
Cuando alguien mide un kernel con tiling y obtiene una mejora de tres veces en lugar de dieciséis, la conclusión habitual es que el tiling “no funciona tan bien en la práctica”. Casi siempre es al revés: el tiling ha hecho exactamente lo que prometía y lo que falla es la línea base. La versión ingenua se beneficia mucho de la caché L2, que retiene la matriz B entera si es pequeña, así que su rendimiento medido está muy por encima del que predice el modelo, y la mejora relativa parece pequeña. Eso tiene una consecuencia práctica importante: el beneficio del tiling crece con el tamaño de la matriz, porque las matrices grandes dejan de caber en caché y la versión ingenua colapsa hacia su predicción teórica. Medir con matrices de 256 y concluir que el tiling aporta poco es un error clásico. El segundo malentendido es el opuesto y más peligroso: creer que como el tiling es una optimización de memoria, se puede aplicar mecánicamente a cualquier kernel. Solo funciona si hay reutilización, y la reutilización de un algoritmo es una propiedad matemática, no una decisión de implementación. En una GEMM cada elemento se usa N veces y el tiling gana. En una suma de vectores cada elemento se usa una vez y el tiling no puede ganar nada, porque no hay nada que reutilizar; el kernel ya lee el mínimo posible y está en su techo. La pregunta que hay que hacerse antes de meter memoria compartida en un kernel es siempre la misma: cuántas veces se lee cada byte. Si la respuesta es uno, cierra el editor.