En octubre de 2023, PyTorch 2.1 añadió soporte para compilar funciones de NumPy a código C++ o CUDA usando Triton, e incluso paralelizarlas con OpenMP cuando es posible.

En el blog de PyTorch enseñaron cómo consiguieron ejecuciones hasta 200x más rápidas para un paso simple de un algoritmo de kmeans.

Después de ver esos resultados, e inspirado por este post, necesitaba probarlo para comprobar si podía conseguir aceleraciones de ese orden en ejemplos de mi propio código.

Probándolo con algoritmos genéticos

Últimamente me han interesado los algoritmos genéticos como forma de búsqueda heurística. Son divertidos de implementar y los resultados pueden ser sorprendentes. Aun así, si no se diseñan bien, los tiempos de ejecución pueden dispararse, especialmente con funciones de fitness complejas, poblaciones grandes y cromosomas sofisticados. Esto me llevó a preguntarme cuánta aceleración podía conseguir compilando código NumPy con este nuevo método.

Ejemplo 1

Para empezar con algo simple, quise medir la aceleración en mi máquina usando una función de fitness sencilla pero computacionalmente intensa: la suma de un array.

Supongamos una población de 900.000 cromosomas, cada uno con 1.000 genes, y que la función de fitness de cada cromosoma es la suma de todos sus genes. Evaluemos los tiempos usando distintas implementaciones.

Primero, la función de fitness en Python puro:

population = [random.choices(range(1, 25), k=1000) for _ in range(N_POPULATION)]

def fitness_fn_python(individual):
	return sum(individual)

fitnesses = [fitness_fn_python(individual) for individual in population]

Ejecutar esto con 900.000 cromosomas tarda 6.91s, sin contar el tiempo de generar la población. Veamos cuánto podemos acelerarlo con numpy y torch.

La versión NumPy sería:

population = np.random.randint(1, 25, size=(N_POPULATION, 1000), dtype=np.int8)

def fitness_fn_numpy(individual):
	return np.sum(individual)

fitnesses = np.apply_along_axis(fitness_fn_numpy, 1, population)

Se ejecuta en 5.19s, un 33% más rápido que Python puro. Aun así sigue siendo lento, porque no está implementado de la forma más eficiente.

Una implementación mejor con NumPy sería:

def fitness_fn_numpy_better(population):
	return np.sum(population, axis=1)

fitnesses = fitness_fn_numpy_better(population)

Esta versión se ejecuta en 0.51s, unas 13.5 veces más rápida que Python puro.

Ahora compilemos esta función con torch.compile:

import torch
fitness_fn_compiled = torch.compile(fitness_fn_numpy_better)

fitnesses = fitness_fn_compiled(population)

La primera ejecución tarda 0.59s por el coste de compilación. En ejecuciones posteriores baja a 0.41s, unas 17 veces más rápido que Python puro.

Para ir todavía más rápido, podemos escribir esta función en PyTorch, compilarla y ejecutarla en CUDA.

Primero, la función de fitness en PyTorch:

population = torch.randint(1, 25, size=(N_POPULATION, 1000), dtype=torch.int8)

def fitness_fn_torch(population):
	return torch.sum(population, dim=1)

fitnesses = fitness_fn_torch(population)

Sin compilar, esta función tarda 2.52s.

Al compilarla:

fitness_fn_compiled = torch.compile(fitness_fn_torch)
fitnesses = fitness_fn_compiled(population)

la primera ejecución, incluyendo compilación, tarda 4.10s. Pero en ejecuciones posteriores baja a 0.067s, que es 103 veces más rápido que Python puro y 9 veces más rápido que la versión NumPy sin compilar.

El último empujón llega al ejecutar esta función compilada en CUDA:

population = population.to('cuda')
with torch.device('cuda'):
    fitnesses = fitness_fn_compiled(population)

Esto consigue un tiempo de 0.055s, unas 125 veces más rápido que Python puro y 10 veces más rápido que NumPy, incluso en una GPU modesta de portátil.

En este ejemplo, el máximo rendimiento se consiguió escribiendo y compilando la función de fitness en PyTorch, con tiempos mejores que usando NumPy compilado por PyTorch.

Ejemplo 2

Este ejemplo se parece más a un problema real de optimización numérica con restricciones. Estos problemas buscan los valores óptimos que minimizan o maximizan una función matemática bajo ciertas condiciones.

Considera esta función:

f(x,y,z)=x2+2y2+3z2+4xy5yzf(x, y, z) = x^2 + 2y^2 + 3z^2 + 4xy - 5yz

con las restricciones:

x13<y7z<0x+y+z3x2+z210\begin{align*} & x \geq 1 \\ & 3 < y \leq 7 \\ & z < 0 \\ & x + y + z \geq 3 \\ & x^2 + z^2 \leq 10 \\ \end{align*}

En Python puro, la función de fitness es:

import random
import math

population = [[random.uniform(-25, 25) for _ in range(3)] for _ in range(N_POPULATION)]

def function_to_optimize(individual):
    x, y, z = individual
    return x**2 + 2*y**2 + 3*z**2 + 4*x*y - 5*y*z

def check_constraints(individual):
    x, y, z = individual
    if x < 1:
        return False
    if not (3 < y <= 7):
        return False
    if z >= 0:
        return False
    if (x + y + z) < 3:
        return False
    if (x**2 + z**2) > 10:
        return False
    return True

def fitness_fn_python(individual):
    if not check_constraints(individual):
        return math.inf
    return -function_to_optimize(individual)

fitnesses = [fitness_fn_python(individual) for individual in population]

Ejecutarlo para 900.000 cromosomas tarda 270ms.

Traduciendo la función a NumPy vectorizado:

import numpy as np

population = np.random.uniform(-25, 25, size=(N_POPULATION, 3))

def function_to_optimize(solutions):
    x = solutions[:, 0]
    y = solutions[:, 1]
    z = solutions[:, 2]
    return x**2 + 2*y**2 + 3*z**2 + 4*x*y - 5*y*z

# Vectorized function to check constraints
def check_constraints(solutions):
    x = solutions[:, 0]
    y = solutions[:, 1]
    z = solutions[:, 2]
    constraints = (x >= 1) & (y > 3) & (y <= 7) & (z < 0) & ((x + y + z) >= 3) & ((x**2 + z**2) <= 10)
    return constraints

# Vectorized fitness function
def fitness_function_numpy(solutions):
    # Check constraints for all solutions
    valid_constraints = check_constraints(solutions)
    # Apply constraints
    fitness = np.where(valid_constraints, -function_to_optimize(solutions), np.inf)
    return fitness

fitnesses = fitness_function_numpy(population)

La versión NumPy vectorizada tarda 48ms, 5.6 veces más rápida que Python puro.

Si compilamos esta función NumPy con torch.compile:

import torch 
population = torch.rand(N_POPULATION, 3) * 50 - 25

check_constraints_compiled = torch.compile(check_constraints)
function_to_optimize_compiled = torch.compile(function_to_optimize)

def fitness_fn_pytorch(population):
    # Check constraints for all solutions
    valid_constraints = check_constraints_compiled(population)
    # Apply constraints
    fitness = np.where(valid_constraints, -function_to_optimize_compiled(population), np.inf)
    return fitness

fitnesses = fitness_fn_pytorch(population)

La función compilada tarda 6.2ms, excluyendo el coste de compilación. Esto es 44 veces más rápido que Python puro y 7 veces más rápido que NumPy estándar.

Al ejecutarla en CUDA, tarda 3.7ms, 73 veces más rápido que Python puro y 13 veces más rápido que NumPy.

Ejemplo 3

De forma parecida al segundo ejemplo, intentamos optimizar esta función:

f(x,y,z)=sin(x)+x22cos(y)+y2+z2f(x, y, z) = \sin(x) + x^2 - 2\cos(y) + y^2 + z^2

sujeta a estas restricciones:

x00yπz5x+z1cos(y)+z0.5\begin{align*} & x \geq 0 \\ & 0 \leq y \leq \pi \\ & z \leq 5 \\ & x + z \geq 1 \\ & \cos(y) + z \geq 0.5 \\ \end{align*}

No incluyo el código de este experimento porque es muy parecido al del segundo ejemplo, pero los tiempos fueron:

  • 240ms en Python puro.
  • 82ms en NumPy, 3x más rápido que Python puro.
  • 36.8ms con NumPy compilado por torch, 6.5x más rápido que Python puro y 2x más rápido que NumPy.
  • 22.6ms con NumPy compilado por torch en CUDA, 10x más rápido que Python puro y 3.6x más rápido que NumPy.

Conclusión

Esta actualización de PyTorch permite conseguir aceleraciones importantes en código numérico de Python sin tener que reescribirlo entero en otro lenguaje.

En estos ejemplos, las mayores mejoras aparecieron en código que ya estaba vectorizado. Compilar la versión de NumPy ayudó, y reescribir la función en PyTorch dio resultados todavía mejores. La aceleración exacta depende de la carga de trabajo, pero los resultados muestran que merece la pena probarlo antes de mover código crítico para el rendimiento a C++, Rust o una implementación CUDA propia.

En cargas de trabajo de machine learning, optimización y análisis de datos, poder probar esto con unos pocos cambios puede ahorrar bastante tiempo. El código se mantiene cerca de la implementación original en Python, mientras que la versión compilada puede aprovechar el paralelismo de CPU o CUDA cuando la carga de trabajo lo permite.

El coste de compilación sigue importando, así que resulta más útil cuando la misma función se ejecuta suficientes veces. Aun así, ofrece a los programadores de Python otra opción para acelerar código numérico sin añadir demasiada complejidad al flujo de trabajo.

Notas

Limitaciones en la paralelización

Aunque esta nueva función de compilación logra detectar partes paralelizables del código, en el momento de escribir el artículo llevaba menos de un mes implementada. A veces no identifica segmentos que podrían distribuirse entre varios cores, especialmente si contienen lógica condicional no expresada con operaciones NumPy, como ocurre con algunas comprobaciones de restricciones.

Es posible que esto mejore en el futuro, pero no es mala idea escribir el código de forma que sea más fácil para TorchInductor si queremos obtener las mejores aceleraciones posibles.

También conviene revisar el código C++/Triton generado usando la variable de entorno TORCH_LOGS=output_code.

El coste de compilación importa

Al usar este método hay que pensar si el coste de compilación merece la pena para un caso de uso concreto. Las funciones de ejemplo tienen tiempos de compilación cortos, pero con tareas más complejas el tiempo necesario para compilar podría comerse el ahorro conseguido durante la ejecución.

En la documentación de PyTorch aparece este ejemplo:

model = init_model()
opt = torch.optim.Adam(model.parameters())

def train(mod, data):
    opt.zero_grad(True)
    pred = mod(data[0])
    loss = torch.nn.CrossEntropyLoss()(pred, data[1])
    loss.backward()
    opt.step()

eager_times = []
for i in range(N_ITERS):
    inp = generate_data(16)
    _, eager_time = timed(lambda: train(model, inp))
    eager_times.append(eager_time)
    print(f"eager train time {i}: {eager_time}")
print("~" * 10)

model = init_model()
opt = torch.optim.Adam(model.parameters())
train_opt = torch.compile(train, mode="reduce-overhead")

compile_times = []
for i in range(N_ITERS):
    inp = generate_data(16)
    _, compile_time = timed(lambda: train_opt(model, inp))
    compile_times.append(compile_time)
    print(f"compile train time {i}: {compile_time}")
print("~" * 10)

eager_med = np.median(eager_times)
compile_med = np.median(compile_times)
speedup = eager_med / compile_med
assert(speedup > 1)
print(f"(train) eager median: {eager_med}, compile median: {compile_med}, speedup: {speedup}x")
print("~" * 10)

La salida es:

eager train time 0: 0.3666875
eager train time 1: 0.06788114929199218
eager train time 2: 0.06595875549316406
eager train time 3: 0.06623709106445312
eager train time 4: 0.06622108459472656
eager train time 5: 0.06640013122558594
eager train time 6: 0.06633692932128907
eager train time 7: 0.06608159637451172
eager train time 8: 0.06633171081542968
eager train time 9: 0.06623030090332031
~~~~~~~~~~
skipping cudagraphs due to input mutation
compile train time 0: 285.17234375
compile train time 1: 3.08613623046875
compile train time 2: 0.04830195236206054
compile train time 3: 0.03755414581298828
compile train time 4: 0.037427520751953124
compile train time 5: 0.03738982391357422
compile train time 6: 0.03739823913574219
compile train time 7: 0.03733385467529297
compile train time 8: 0.03738035202026367
compile train time 9: 0.03721017456054688
~~~~~~~~~~

Los tiempos de entrenamiento son 1.77 veces más rápidos, pero la primera iteración tarda 790 veces más. Por tanto, compilar solo merece la pena si hay suficientes pasos de entrenamiento.

Referencias

Making Python 100x faster with less than 100 lines of Rust - Ohad Ravid

torch.compile Tutorial - William Wen

Compiling NumPy code into C++ or CUDA via torch.compile - PyTorch blog