Функция compile
Функция compile преобразует модель PyTorch в
оптимизированную версию, которая выполняется быстрее
за счет компиляции графа вычислений. Компиляция
особенно эффективна для моделей, которые многократно
выполняются с одинаковой структурой тензоров.
Параметры функции:
-
model- модель PyTorch для оптимизации -
fullgraph- булево значение, указывающее, должен ли весь граф быть скомпилирован как единое целое -
dynamic- булево значение, позволяющее компилировать модель с поддержкой изменяющихся размеров тензоров
Синтаксис
torch.compile(model, [fullgraph], [dynamic])
Пример
Давайте создадим простую модель линейных слоев и скомпилируем ее для ускорения на GPU:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(10, 20)
self.fc2 = nn.Linear(20, 1)
def forward(self, x):
x = torch.relu(self.fc1(x))
return self.fc2(x)
model = SimpleModel()
x = torch.randn(100, 10).cuda()
compiled_model = torch.compile(model)
res = compiled_model(x)
Результат выполнения кода:
tensor([[-0.0717],
[ 0.1786],
[ 0.0224],
...,
[ 0.1834],
[ 0.0093],
[ 0.1693]], device='cuda:0', grad_fn=<ViewBackward0>)
Пример
Давайте скомпилируем модель с поддержкой динамических
размеров тензоров с помощью параметра dynamic:
import torch
import torch.nn as nn
class DynamicModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
model = DynamicModel().cuda()
compiled_model = torch.compile(model, dynamic=True)
x1 = torch.randn(32, 10).cuda()
x2 = torch.randn(64, 10).cuda()
res1 = compiled_model(x1)
res2 = compiled_model(x2)
print(res1.shape, res2.shape)
Результат выполнения кода:
torch.Size([32, 5]) torch.Size([64, 5])
Пример
Давайте укажем параметр fullgraph для компиляции
всего графа вычислений целиком:
import torch
model = torch.nn.Linear(10, 5).cuda()
compiled_model = torch.compile(model, fullgraph=True)
x = torch.randn(100, 10).cuda()
res = compiled_model(x)
print(res.shape)
Результат выполнения кода:
torch.Size([100, 5])
Смотрите также
-
функцию
is_available,
которая проверяет доступность CUDA -
функцию
synchronize,
которая синхронизирует все операции на GPU -
функцию
set_device,
которая устанавливает устройство для вычислений -
функцию
empty_cache,
которая очищает кэш памяти GPU