Radix sort paralelo sobre el scan
Ordenación en tiempo lineal sin comparar: histograma por bloque, scan global de los conteos y dispersión estable con ranking local.
El radix sort no compara elementos. Mira un puñado de bits de la clave, cuenta cuántos elementos caen en cada valor posible de esos bits, y usa un scan para saber exactamente dónde va cada uno. Repitiendo eso desde los bits bajos hacia los altos, con una partición estable en cada pasada, el array queda ordenado. Su coste es lineal en el número de elementos y su número de pasadas depende solo de los bits de la clave, así que a partir de cierto tamaño gana a cualquier red de ordenación por un margen que crece. Y todo él se apoya en el scan del nivel anterior: si aquello estaba claro, esto es montaje.
- Justificar por qué el radix sort necesita que cada pasada sea estable.
- Construir el histograma por bloque con la disposición que hace útil el scan.
- Implementar el ranking local con particiones sucesivas de un bit.
- Convertir claves con signo y claves en coma flotante en enteros ordenables.
La idea y el requisito de estabilidad
Ordenar por el dígito menos significativo primero solo funciona si cada pasada conserva el orden relativo de los elementos que tienen el mismo dígito. Si no, la pasada por el dígito de las decenas destruiría el orden que estableció la de las unidades.
entrada 170 045 075 090 802 024 002 066
por unidades 170 090 802 002 024 045 075 066
por decenas 802 002 024 045 066 170 075 090
por centenas 002 024 045 066 075 090 170 802
Fíjate en la segunda línea: 802 y 002 tienen el mismo dígito de unidades y 802 venía antes, así que sigue antes. Si esa pasada los hubiera intercambiado, la tercera línea saldría mal. La estabilidad no es una propiedad deseable del radix sort, es su condición de corrección.
Con r bits por pasada hay 2^r cubos y hacen falta 32/r pasadas para una clave de 32 bits. Con r = 4 son 16 cubos y 8 pasadas. Subir a r = 8 daría 4 pasadas pero 256 cubos, y el ranking local dentro del bloque se vuelve mucho más caro; cuatro bits es el punto dulce que usa casi todo el mundo.
Las tres fases de una pasada
Una pasada se descompone en tres kernels, y el del medio es el scan de varios bloques sin modificar.
Histograma. Cada bloque de B elementos cuenta cuántos tiene de cada dígito y escribe 16 números. La disposición del array de conteos es la decisión importante: se indexa por digito · numBloques + bloque, no al revés.
const WG: u32 = 256u;
const CUBOS: u32 = 16u;
struct Params { conteo: u32, bitBase: u32, numBloques: u32 };
@group(0) @binding(0) var<uniform> p: Params;
@group(0) @binding(1) var<storage, read> claves: array<u32>;
@group(0) @binding(2) var<storage, read_write> conteos: array<u32>;
var<workgroup> local: array<atomic<u32>, CUBOS>;
@compute @workgroup_size(WG)
fn histograma(@builtin(global_invocation_id) gid: vec3u,
@builtin(local_invocation_index) t: u32,
@builtin(workgroup_id) wid: vec3u) {
if (t < CUBOS) { atomicStore(&local[t], 0u); }
workgroupBarrier();
let i = gid.x;
if (i < p.conteo) {
let d = (claves[i] >> p.bitBase) & (CUBOS - 1u);
atomicAdd(&local[d], 1u);
}
workgroupBarrier();
if (t < CUBOS) {
// Disposicion por digito: todos los bloques de un digito, contiguos.
// Cada ranura la escribe una sola invocacion, asi que no es atomica.
conteos[t * p.numBloques + wid.x] = atomicLoad(&local[t]);
}
}
Aquí el if (i < p.conteo) sí es legal aunque haya barreras, porque no es un return: las invocaciones sobrantes se saltan el incremento pero siguen llegando a todas las barreras.
Esa disposición por dígito es lo que hace que un único scan exclusivo del array entero de conteos produzca lo que hace falta. La posición d · numBloques + b del scan contiene la suma de todos los conteos anteriores: todos los elementos de dígitos menores que d, de todos los bloques, más los elementos de dígito d de los bloques anteriores a b. Eso es exactamente la posición global donde el bloque b debe empezar a escribir sus elementos de dígito d. Con la disposición traspuesta el scan no daría nada útil.
Dispersión. Cada bloque vuelve a leer sus elementos, los ordena localmente por el dígito, calcula el rango de cada uno dentro de su dígito, y escribe en el destino global.
El ranking local con particiones de un bit
La forma correcta y barata de ordenar un bloque por un dígito de cuatro bits es hacer cuatro particiones estables de un bit, cada una construida sobre un scan local. Una partición de un bit separa los ceros de los unos conservando el orden dentro de cada grupo.
var<workgroup> claveL: array<u32, WG>;
var<workgroup> cargaL: array<u32, WG>;
var<workgroup> sBuf: array<u32, 2u * WG>;
var<workgroup> sTotal: u32;
// Scan exclusivo de un valor por invocacion. Deja el total en sTotal.
fn scanExclusivo(t: u32, v: u32) -> u32 {
var lee: u32 = 0u;
var esc: u32 = WG;
sBuf[lee + t] = v;
workgroupBarrier();
for (var d: u32 = 1u; d < WG; d <<= 1u) {
var r = sBuf[lee + t];
if (t >= d) { r = r + sBuf[lee + t - d]; }
sBuf[esc + t] = r;
workgroupBarrier();
let x = lee; lee = esc; esc = x;
}
if (t == WG - 1u) { sTotal = sBuf[lee + t]; }
var e: u32 = 0u;
if (t > 0u) { e = sBuf[lee + t - 1u]; }
workgroupBarrier();
return e;
}
// Particion estable del bloque por el bit b.
fn particion(t: u32, b: u32) {
let k = claveL[t];
let c = cargaL[t];
let uno = (k >> b) & 1u;
let prefCeros = scanExclusivo(t, 1u - uno);
let totalCeros = sTotal;
var destino = prefCeros;
if (uno == 1u) { destino = totalCeros + (t - prefCeros); }
claveL[destino] = k;
cargaL[destino] = c;
workgroupBarrier();
}
La escritura en claveL[destino] es segura sin barrera previa porque scanExclusivo ya ha ejecutado varias barreras después de que todas las invocaciones leyeran su k y su c.
Y el kernel de dispersión completo:
@group(0) @binding(3) var<storage, read> desplaz: array<u32>; // scan de conteos
@group(0) @binding(4) var<storage, read_write> destClaves: array<u32>;
@group(0) @binding(5) var<storage, read_write> destCargas: array<u32>;
@group(0) @binding(6) var<storage, read> cargas: array<u32>;
var<workgroup> inicio: array<u32, CUBOS>;
@compute @workgroup_size(WG)
fn dispersar(@builtin(global_invocation_id) gid: vec3u,
@builtin(local_invocation_index) t: u32,
@builtin(workgroup_id) wid: vec3u) {
let i = gid.x;
// Relleno con el centinela maximo: cae siempre en el ultimo cubo y
// detras de los elementos reales, asi que no altera sus rangos.
var k: u32 = 0xffffffffu;
var c: u32 = 0u;
if (i < p.conteo) { k = claves[i]; c = cargas[i]; }
claveL[t] = k;
cargaL[t] = c;
if (t < CUBOS) { inicio[t] = 0u; }
workgroupBarrier();
// Cuatro particiones de un bit: el bloque queda ordenado por el digito.
particion(t, p.bitBase + 0u);
particion(t, p.bitBase + 1u);
particion(t, p.bitBase + 2u);
particion(t, p.bitBase + 3u);
// Primera posicion local de cada digito. El bloque ya esta ordenado,
// asi que basta detectar el cambio respecto al vecino anterior.
let d = (claveL[t] >> p.bitBase) & (CUBOS - 1u);
if (t > 0u) {
let dPrev = (claveL[t - 1u] >> p.bitBase) & (CUBOS - 1u);
if (d != dPrev) { inicio[d] = t; }
}
workgroupBarrier();
let rango = t - inicio[d];
let base = desplaz[d * p.numBloques + wid.x];
if (claveL[t] != 0xffffffffu || i < p.conteo) {
destClaves[base + rango] = claveL[t];
destCargas[base + rango] = cargaL[t];
}
}
La condición final merece un comentario: escribe si el elemento no es relleno. Como el relleno usa el centinela máximo, un elemento real que valga exactamente 0xffffffff también se descartaría; por eso se acepta si el índice original estaba dentro del conteo. Si tus claves nunca alcanzan el máximo, la comprobación se simplifica a una sola condición.
Cada pasada intercambia origen y destino: ocho pasadas con dos buffers en ping-pong dejan el resultado en el original, lo cual es cómodo.
const PASADAS = 8;
for (let pasada = 0; pasada < PASADAS; pasada++) {
const bitBase = pasada * 4;
device.queue.writeBuffer(uni[pasada], 0,
new Uint32Array([n, bitBase, numBloques]));
// histograma -> scan (3 dispatches) -> dispersion
// los bind groups alternan A->B y B->A segun la paridad
}
Salen cinco dispatches por pasada, cuarenta en total para claves de 32 bits, frente a los 78 de un bitonic sort del mismo tamaño. Y el trabajo es lineal, no O(n·log²n).
Claves con signo y en coma flotante
El radix sort ordena enteros sin signo. Cualquier otra clave hay que transformarla en un u32 que preserve el orden, y deshacer la transformación al final.
Para enteros con signo basta invertir el bit de signo, porque en complemento a dos los negativos tienen ese bit a uno y ordenan por debajo:
fn claveDeI32(v: i32) -> u32 { return bitcast<u32>(v) ^ 0x80000000u; }
fn i32DeClave(k: u32) -> i32 { return bitcast<i32>(k ^ 0x80000000u); }
Para f32 hay que tener en cuenta además que los negativos ordenan al revés como patrón de bits: cuanto mayor la magnitud, mayor el entero. La transformación estándar invierte todos los bits de los negativos y solo el de signo de los positivos:
fn claveDeF32(v: f32) -> u32 {
let u = bitcast<u32>(v);
// Si el signo es 1 -> invertir todo. Si es 0 -> invertir solo el signo.
let mascara = select(0x80000000u, 0xffffffffu, (u >> 31u) == 1u);
return u ^ mascara;
}
fn f32DeClave(k: u32) -> f32 {
let mascara = select(0xffffffffu, 0x80000000u, (k >> 31u) == 1u);
return bitcast<f32>(k ^ mascara);
}
Ese par de funciones es correcto para todos los flotantes finitos y para los infinitos. Los NaN quedan ordenados en un extremo, que es lo que uno quiere si de todas formas son un error. El menos cero y el más cero se separan por un bit, con el menos cero primero, cosa que casi nunca importa y conviene saber.
Cuando la clave tiene menos de 32 bits significativos —una profundidad cuantizada a 16, un índice de celda de una rejilla de 22 bits— el número de pasadas baja proporcionalmente. Un índice de celda de 16 bits se ordena en cuatro pasadas en vez de ocho, y eso divide el tiempo por dos.
El argumento de venta del radix sort es su coste lineal, y es cierto en operaciones. En tiempo medido, lo que domina es otra cosa: la escritura de la fase de dispersión es dispersa. Los 256 elementos de un bloque van a 16 destinos distintos, así que un warp de 32 carriles escribe en varias regiones alejadas de VRAM y la transacción de memoria se fragmenta. Una escritura coalescente de 32 valores consecutivos es una sola transacción; 32 valores repartidos en 16 zonas son hasta 16. Ese factor es la razón por la que el ranking local no es una optimización opcional sino el corazón de la implementación: al ordenar el bloque localmente antes de escribir, los elementos que van al mismo destino quedan contiguos, y la escritura pasa de dispersa a un puñado de tramos consecutivos. La diferencia entre la versión con ranking local y la ingenua que escribe cada elemento donde le toca ronda un factor de tres o cuatro. Y explica una decisión que sorprende al leer implementaciones reales: por qué se gastan cuatro scans locales completos, con sus treinta y tantas barreras, en reordenar un bloque que después se va a escribir de todas formas. La respuesta es que esas barreras cuestan nanosegundos en memoria compartida y las transacciones de VRAM que ahorran cuestan cientos de ciclos cada una. Es el mismo principio que hace ganar al tiling: en una GPU, casi cualquier cantidad de trabajo en memoria compartida sale más barata que un acceso mal formado a memoria global.