В этом документе описывается, как писать и использовать пользовательские вызовы XLA с помощью библиотеки XLA FFI. Пользовательский вызов — это механизм для описания внешней «операции» в модуле HLO компилятору XLA (во время компиляции), а XLA FFI — это механизм для регистрации реализации таких операций в XLA (во время выполнения). FFI расшифровывается как «интерфейс внешних функций» и представляет собой набор C API, определяющих бинарный интерфейс (ABI) для вызова XLA из внешнего кода, написанного на других языках программирования. XLA предоставляет только заголовочные привязки для XLA FFI, написанного на C++, что скрывает от конечного пользователя все низкоуровневые детали базовых C API.
Пользовательские звонки JAX + XLA
Примеры комплексной интеграции пользовательских вызовов и XLA FFI с JAX см. в документации JAX .
Привязка XLA FFI
XLA FFI-привязка — это спецификация пользовательской сигнатуры вызова, определяемая на этапе компиляции: пользовательские аргументы вызова, атрибуты и их типы, а также дополнительные параметры, передаваемые через контекст выполнения (например, поток GPU для бэкенда GPU). XLA FFI-привязка может быть привязана к любому вызываемому объекту C++ (указатель на функцию, лямбда-функция и т. д.) с совместимой сигнатурой operator() . Созданный обработчик декодирует кадр вызова XLA FFI (определенный стабильным API C), проверяет тип всех параметров и передает декодированные результаты в определяемую пользователем функцию обратного вызова.
В основе привязки XLA FFI лежит метапрограммирование на основе шаблонов, позволяющее компилировать созданный обработчик в наиболее эффективный машинный код. Накладные расходы во время выполнения составляют порядка нескольких наносекунд для каждого пользовательского параметра вызова.
Точки настройки XLA FFI реализованы в виде специализаций шаблонов, и пользователи могут определять, как декодировать свои пользовательские типы, то есть можно определить пользовательское декодирование для типов enum class , заданных пользователем.
Возвращение ошибок из пользовательских вызовов
Реализации пользовательских вызовов должны возвращать значение xla::ffi::Error , чтобы сигнализировать об успехе или ошибке среде выполнения XLA. Это похоже на absl::Status и имеет тот же набор кодов ошибок. Мы не используем absl::Status потому что у него нет стабильного ABI, и передавать его между динамически загружаемой библиотекой пользовательских вызовов и самой XLA было бы небезопасно.
// Handler that always returns an error.
auto always_error = Ffi::Bind().To(
[]() { return Error(ErrorCode::kInternal, "Oops!"); });
// Handler that always returns a success.
auto always_success = Ffi::Bind().To(
[]() { return Error::Success(); });
Аргументы и результаты буферизации
XLA использует стиль передачи результатов в целевые хранилища: пользовательские вызовы (или любые другие операции XLA) не выделяют память для результатов, а вместо этого записывают данные в целевые хранилища, переданные средой выполнения XLA. XLA использует статическое назначение буферов и выделяет буферы для всех значений на основе их диапазонов значений во время компиляции.
Результаты, передаваемые обработчикам FFI, обернуты в шаблон Result<T> , имеющий семантику указателя: operator-> обеспечивает доступ к базовому параметру.
Функция AnyBuffer предоставляет доступ к параметрам пользовательского буфера вызова любого типа данных. Это полезно, когда пользовательский вызов имеет универсальную реализацию, работающую с несколькими типами данных, и реализация пользовательского вызова осуществляет диспетчеризацию во время выполнения в зависимости от типа данных. AnyBuffer предоставляет доступ к типу данных буфера, его размерам и указателю на сам буфер.
%0 = "stablehlo.custom_call"(%arg0) {
call_target_name = "foo",
api_version = 4 : i32
} : (tensor<2x2xf32>) -> tensor<2x2xf32>
// Buffers of any number of dimensions and data type.
auto handler = Ffi::Bind().Arg<AnyBuffer>().Ret<AnyBuffer>().To(
[](AnyBuffer arg, Result<AnyBuffer> res) -> Error {
void* arg_data = arg.untyped_data();
void* res_data = res->untyped_data();
return Error::Success();
});
Аргументы и результаты ограниченного буфера
Buffer позволяет добавлять ограничения на тип данных буфера и количество измерений, которые будут автоматически проверяться обработчиком и возвращать ошибку в среду выполнения XLA, если аргументы среды выполнения не соответствуют сигнатуре обработчика FFI.
// Buffers of any number of dimensions and F32 data type.
auto handler = Ffi::Bind().Arg<Buffer<F32>>().Ret<Buffer<F32>>().To(
[](Buffer<F32> arg, Result<Buffer<F32>> res) -> Error {
float* arg_data = arg.typed_data();
float* res_data = res->typed_data();
return Error::Success();
});
// Buffers of number of dimensions 2 and F32 data type.
auto handler = Ffi::Bind().Arg<BufferR2<F32>>().Ret<BufferR2<F32>>().To(
[](BufferR2<F32> arg, Result<BufferR2<F32>> res) -> Error {
float* arg_data = arg.typed_data();
float* res_data = res->typed_data();
return Error::Success();
});
Сопоставление и проверка буферов
Шаблоны буферов позволяют выражать ограничения, которые являются более специфичными, чем типы аргументов и результатов в привязке FFI. Они полезны для уточнения AnyBuffer до конкретного типа буфера, проверки размеров и проверки взаимосвязей между несколькими формами буферов.
Шаблоны являются неизменяемыми и могут указывать тип данных и ранг либо напрямую, либо с помощью соответствующих модификаторов:
namespace m = ::xla::ffi::match;
m::Buffer<F32, 2>();
m::Buffer().WithDType<F32>().WithRank<2>();
Функция Match проверяет шаблон и возвращает максимально ограниченный тип буфера. Для уточнения AnyBuffer шаблон должен указывать ровно один тип данных и один ранг. Сопоставление уже типизированного буфера проверяет любые дополнительные ограничения и сохраняет его тип.
namespace m = ::xla::ffi::match;
auto handler = Ffi::Bind().Arg<AnyBuffer>().Ret<AnyBuffer>().To(
[](AnyBuffer input, Result<AnyBuffer> output) -> Error {
int64_t rows;
int64_t cols;
ErrorOr<BufferR2<F32>> typed_input =
Match("input", input,
m::Buffer<F32>().WithDims(&rows, &cols));
if (typed_input.has_error()) {
return std::move(typed_input).error();
}
ErrorOr<Result<BufferR2<F32>>> typed_output =
Match("output", output,
m::Buffer<F32>().WithDims(rows, cols));
if (typed_output.has_error()) {
return std::move(typed_output).error();
}
float* input_data = typed_input->typed_data();
float* output_data = (*typed_output)->typed_data();
return Error::Success();
});
В этом примере input соответствующие шаблону, фиксируют два его измерения, а output проверяют, имеют ли они ту же форму. Фиксация данных происходит только тогда, когда весь шаблон успешно совпадает.
Шаблон измерения может быть неограниченным, фиксированным или захватываемым. Фиксированные значения и указатели захвата неявно преобразуются в шаблоны измерений:
m::Dim() // Any dimension size.
m::Dim(16) // A dimension whose size is 16.
m::Dim(&rows) // Any dimension size, captured in `rows` on success.
m::Buffer<F32>().WithDims(&rows, 16);
WithDims описывает полную форму и задает ее ранг на основе количества аргументов. WithDim<I> или WithDim(index, ...) ограничивает размеры позиционно, не фиксируя полный ранг. Позиционное ограничение требует существования указанного размера, но оставляет все остальные размеры без ограничений:
ErrorOr<BufferR4<F32>> MatchActivations(AnyBuffer input) {
return Match(
"input", input,
m::Buffer<F32, 4>().WithDim<0>(1).WithDim<3>(128));
}
Error VerifyDimension(AnyBuffer input, size_t index, int64_t size) {
return Verify("input", input, m::Buffer().WithDim(index, size));
}
MatchActivations требует буфер F32 ранга 4 с размерами 0 и 3, равными 1 и 128 соответственно; размеры 1 и 2 не ограничены. Без явного WithRank функция WithDim(index, ...) принимает любой ранг больше index .
В параметрах измерения также можно отображать взаимосвязи между позициями:
auto matrix = m::Buffer<F32>().WithDims(&rows, 128);
int64_t n;
ErrorOr<BufferR2<F32>> square =
Match("matrix", buffer,
m::Buffer<F32>().WithDims(&n, &n));
Когда один и тот же указатель захвата встречается несколько раз, все соответствующие измерения должны иметь одинаковый размер. Поэтому вызов Match , описанный выше, принимает только квадратные матрицы. Захват измерений записывается только после успешного выполнения всего шаблона и остается неизменным в случае неудачи.
Для сопоставления буферов целиком не требуется учитывать каждое измерение. WithShapeOf сопоставляет полную форму другого буфера во время выполнения, допуская при этом другой тип данных. Like также требует одинакового типа данных:
Error VerifySortBuffers(AnyBuffer keys, AnyBuffer values,
Result<AnyBuffer> keys_output) {
Error error =
Verify("values", values, m::Buffer().WithShapeOf(keys));
if (error.failure()) {
return error;
}
return Verify("keys output", *keys_output, m::Buffer().Like(keys));
}
Оба модификатора копируют метаданные буфера-ссылки при построении шаблона; шаблон не сохраняет ссылку на буфер. Поскольку это ограничения времени выполнения, они не уточняют тип возвращаемого значения AnyBuffer . Они в основном полезны с Verify или с Match , когда входной буфер уже имеет конкретный тип.
Используйте Verify , когда буфер должен сохранять свой существующий тип или когда шаблон принимает более одного типа данных или ранга. Например, индексный буфер может принимать либо S32 , либо S64 и обрабатывать данные по их типу после проверки:
namespace m = ::xla::ffi::match;
Error VerifyIndices(AnyBuffer indices) {
auto pattern = m::Buffer().WithDType<S32, S64>().WithRank<1, 2>();
Error error = Verify("indices", indices, pattern);
if (error.failure()) {
return error;
}
// Dispatch an implementation based on `indices.element_type()`.
return Error::Success();
}
Поскольку этот шаблон допускает несколько конкретных типов буферов, он не может уточнить AnyBuffer с помощью Match : уникального возвращаемого типа не существует. Тем не менее, его можно передать в Match с уже типизированным буфером, возвращаемый тип которого уже известен. Уже типизированный буфер также может быть проверен; это обычно полезно для проверки взаимосвязей форм:
Error VerifySameShape(BufferR2<F32> lhs, BufferR2<F32> rhs) {
int64_t rows;
int64_t cols;
Error error = Verify(
"lhs", lhs,
m::Buffer().WithDims(&rows, &cols));
if (error.failure()) {
return error;
}
return Verify("rhs", rhs,
m::Buffer().WithDims(rows, cols));
}
Имя, передаваемое в качестве первого аргумента функциям Match и Verify является обязательным и включается в сообщения об ошибках вместе с описанием несоответствующего ограничения.
Внешний API FFI в xla/ffi/api/ffi.h возвращает Error из Verify и ErrorOr<T> из Match , как в примерах выше. Внутренний API FFI в xla/ffi/ffi.h предоставляет тот же интерфейс сопоставления, используя absl::Status и absl::StatusOr<T> :
ABSL_ASSIGN_OR_RETURN(BufferR2<F32> input,
Match("input", buffer, m::Buffer<F32, 2>()));
ABSL_RETURN_IF_ERROR(Verify(
"input", input,
m::Buffer().WithDims(expected_rows, m::Dim())));
Различные аргументы и результаты
Если количество аргументов и результат могут отличаться в разных случаях пользовательского вызова, их можно декодировать во время выполнения с помощью RemainingArgs и RemainingRets .
auto handler = Ffi::Bind().RemainingArgs().RemainingRets().To(
[](RemainingArgs args, RemainingRets results) -> Error {
ErrorOr<AnyBuffer> arg = args.get<AnyBuffer>(0);
ErrorOr<Result<AnyBuffer>> res = results.get<AnyBuffer>(0);
if (!arg.has_value()) {
return Error(ErrorCode::kInternal, arg.error());
}
if (!res.has_value()) {
return Error(ErrorCode::kInternal, res.error());
}
return Error::Success();
});
Вариативные аргументы и результаты могут быть объявлены после обычных аргументов и результатов, однако связывание обычных аргументов и результатов после вариативных является недопустимым.
auto handler =
Ffi::Bind()
.Arg<AnyBuffer>()
.RemainingArgs()
.Ret<AnyBuffer>()
.RemainingRets()
.To([](AnyBuffer arg, RemainingArgs args, AnyBuffer ret,
RemainingRets results) -> Error { return Error::Success(); });
Атрибуты
XLA FFI поддерживает автоматическое декодирование mlir::DictionaryAttr передаваемого в качестве параметра custom_call backend_config в аргументы обработчика FFI.
%0 = "stablehlo.custom_call"(%arg0) {
call_target_name = "foo",
backend_config= {
i32 = 42 : i32,
str = "string"
},
api_version = 4 : i32
} : (tensor<f32>) -> tensor<f32>
В этом примере пользовательский вызов имеет один аргумент буфера и два атрибута, и XLA FFI может автоматически декодировать их и передать определяемой пользователем вызываемой функции.
auto handler = Ffi::Bind()
.Arg<BufferR0<F32>>()
.Attr<int32_t>("i32")
.Attr<std::string_view>("str")
.To([](BufferR0<F32> buffer, int32_t i32, std::string_view str) {
return Error::Success();
});
Атрибуты перечисления, определяемые пользователем
XLA FFI может автоматически декодировать целочисленные атрибуты MLIR в определяемые пользователем перечисления. Класс перечисления должен иметь тот же базовый целочисленный тип, а декодирование должно быть явно зарегистрировано в XLA FFI.
%0 = "stablehlo.custom_call"(%arg0) {
call_target_name = "foo",
backend_config= {
command = 0 : i32
},
api_version = 4 : i32
} : (tensor<f32>) -> tensor<f32>
enum class Command : int32_t {
kAdd = 0,
kMul = 1,
};
XLA_FFI_REGISTER_ENUM_ATTR_DECODING(Command);
auto handler = Ffi::Bind().Attr<Command>("command").To(
[](Command command) -> Error { return Error::Success(); });
Привязка всех пользовательских атрибутов вызова
Можно получить доступ ко всем пользовательским атрибутам вызова в виде словаря и отложенно декодировать только те атрибуты, которые необходимы во время выполнения.
auto handler = Ffi::Bind().Attrs().To([](Dictionary attrs) -> Error {
ErrorOr<int32_t> i32 = attrs.get<int32_t>("i32");
return Error::Success();
});
Пользовательские структурные атрибуты
XLA FFI может декодировать атрибуты словаря в определяемые пользователем структуры.
%0 = "stablehlo.custom_call"(%arg0) {
call_target_name = "foo",
backend_config= {
range = { lo = 0 : i64, hi = 42 : i64 }
},
api_version = 4 : i32
} : (tensor<f32>) -> tensor<f32>
В приведенном выше примере range — это атрибут mlir::DictionaryAttr , и вместо доступа к полям словаря по имени, его можно автоматически декодировать как структуру C++. Декодирование необходимо явно зарегистрировать с помощью макроса XLA_FFI_REGISTER_STRUCT_ATTR_DECODING (за кулисами он определяет специализацию шаблона в пространстве имен ::xla::ffi , поэтому макрос необходимо добавить в глобальное пространство имен).
struct Range {
int64_t lo;
int64_t hi;
};
XLA_FFI_REGISTER_STRUCT_ATTR_DECODING(Range, StructMember<int64_t>("lo"),
StructMember<int64_t>("hi"));
auto handler = Ffi::Bind().Attr<Range>("range").To([](Range range) -> Error{
return Error::Success();
});
Пользовательские атрибуты можно загружать из словаря, как и любые другие атрибуты. В приведенном ниже примере все пользовательские атрибуты вызова декодируются как Dictionary , а доступ к range можно получить по имени.
auto handler = Ffi::Bind().Attrs().To([](Dictionary attrs) -> Error {
ErrorOr<Range> range = attrs.get<Range>("range");
return Error::Success();
});
Создать пользовательский вызов на ЦП
Вы можете создать инструкцию HLO, представляющую пользовательский вызов, через клиентский API XLA. Например, следующий код использует пользовательский вызов для вычисления A[i] = B[i % 128]+ C[i] на ЦП. (Конечно, вы могли бы — и должны! — сделать это с помощью обычной инструкции HLO.)
#include "xla/client/xla_builder.h"
#include "xla/service/custom_call_target_registry.h"
void do_it() {
xla::XlaBuilder b("do_it");
xla::XlaOp param0 =
xla::Parameter(&b, 0, xla::ShapeUtil::MakeShape(xla::F32, {128}), "p0");
xla::XlaOp param1 =
xla::Parameter(&b, 1, xla::ShapeUtil::MakeShape(xla::F32, {2048}), "p1");
xla::XlaOp custom_call =
xla::CustomCall(&b, "do_custom_call", /*operands=*/{param0, param1},
/*shape=*/xla::ShapeUtil::MakeShape(xla::F32, {2048}),
/*opaque=*/"", /*has_side_effect=*/false,
/*output_operand_aliasing=*/{}, /*literal=*/nullptr,
/*schedule=*/CustomCallSchedule::SCHEDULE_NONE,
/*api_version=*/CustomCallApiVersion::API_VERSION_TYPED_FFI);
}
// Constrain custom call arguments to 1-dimensional buffers of F32 data type.
using BufferF32 = xla::ffi::BufferR1<xla::ffi::DataType::F32>;
// Implement a custom call as a C++ function. Note that we can use `Buffer` type
// defined by XLA FFI that gives us access to buffer data type and shape.
xla::ffi::Error do_custom_call(BufferF32 in0, BufferF32 in1,
xla::ffi::Result<BufferF32> out) {
size_t d0 = in0.dimensions[0];
size_t d1 = in1.dimensions[0];
// Check that dimensions are compatible.
assert(out->dimensions[0] == d1 && "unexpected dimensions");
for (size_t i = 0; i < d1; ++i) {
out->data[i] = in0.data[i % d0] + in1.data[i];
}
}
// Explicitly define an XLA FFI handler signature and bind it to the
// `do_custom_call` implementation. XLA FFI handler can automatically infer
// type signature from the custom call function, but it relies on magical
// template metaprogramming an explicit binding provides and extra level of
// type checking and clearly states custom call author intentions.
XLA_FFI_DEFINE_HANDLER(handler, do_custom_call,
ffi::Ffi::Bind()
.Arg<Buffer>()
.Arg<Buffer>()
.Ret<Buffer>());
// Registers `handler` with and XLA FFI on a "Host" platform.
XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "do_custom_call",
"Host", handler);
Создать пользовательский вызов на графическом процессоре
Регистрация пользовательских вызовов GPU с использованием XLA FFI практически идентична, единственное отличие заключается в том, что для GPU необходимо запросить поток базовой платформы (поток CUDA или ROCM), чтобы иметь возможность запустить ядро на устройстве. Вот пример CUDA, который выполняет те же вычисления ( A[i] = B[i % 128] + C[i] ), что и код для CPU, приведенный выше.
void do_it() { /* same implementation as above */ }
__global__ custom_call_kernel(const float* in0, const float* in1, float* out) {
size_t idx = blockIdx.x * blockDim.x + threadIdx.x;
out[idx] = in0[idx % 128] + in1[idx];
}
void do_custom_call(CUstream stream, BufferF32 in0, BufferF32 in1,
xla::ffi::Result<BufferF32> out) {
size_t d0 = in0.dimensions[0];
size_t d1 = in1.dimensions[0];
size_t d2 = out->dimensions[0];
assert(d0 == 128 && d1 == 2048 && d2 == 2048 && "unexpected dimensions");
const int64_t block_dim = 64;
const int64_t grid_dim = 2048 / block_dim;
custom_call_kernel<<<grid_dim, block_dim, 0, stream>>>(
in0.data, in1.data, out->data);
}
XLA_FFI_DEFINE_HANDLER(handler, do_custom_call,
ffi::Ffi::Bind()
.Ctx<xla::ffi::PlatformStream<CUstream>>()
.Arg<BufferF32>()
.Arg<BufferF32>()
.Ret<BufferF32>());
XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "do_custom_call",
"CUDA", handler);
Прежде всего, обратите внимание, что функция пользовательского вызова GPU по-прежнему является функцией, выполняемой на CPU . Функция do_custom_call CPU отвечает за постановку задач в очередь на GPU. Здесь она запускает ядро CUDA, но могла бы также выполнять другие действия, например, вызывать cuBLAS.
Аргументы и результаты также хранятся на хосте, а член данных содержит указатель на память устройства (например, графического процессора). Буферы, передаваемые в пользовательский обработчик вызовов, имеют форму базовых буферов устройства, поэтому пользовательский вызов может вычислить параметры запуска ядра на их основе.
Передача кортежей в пользовательские вызовы
Рассмотрим следующий пользовательский вызов.
using xla::ShapeUtil;
using xla::F32;
Shape p0_shape = ShapeUtil::MakeTuple({
ShapeUtil::MakeShape(F32, {32}),
ShapeUtil::MakeTuple({
ShapeUtil::MakeShape(F32, {64}),
ShapeUtil::MakeShape(F32, {128}),
}),
ShapeUtil::MakeShape(F32, {256}),
});
xla::XlaOp p0 = xla::Parameter(0, p0_shape, "p0");
Shape out_shape = ShapeUtil::MakeTuple({
ShapeUtil::MakeShape(F32, {512}),
ShapeUtil::MakeShape(F32, {1024}),
});
xla::CustomCall(&b, "do_custom_call", /*operands=*/{p0}, out_shape, ...);
Как на центральном, так и на графическом процессоре кортеж представляется в памяти в виде массива указателей. Когда XLA вызывает пользовательские функции с аргументами или результатами в виде кортежей, он преобразует их в одномерный массив и передает в качестве обычных аргументов или результатов буфера.
Кортежи выводятся в виде временных буферов.
Использование кортежей в качестве входных данных для пользовательских вызовов — это удобство, но оно не является строго необходимым. Если бы мы не поддерживали кортежи в качестве входных данных для пользовательских вызовов, вы всегда могли бы распаковать кортежи с помощью функции get-tuple-element перед передачей их в пользовательский вызов.
С другой стороны, вывод в виде кортежей позволяет делать то, что иначе было бы невозможно.
Очевидная причина использования кортежей в качестве выходных данных заключается в том, что именно с помощью кортежей пользовательский вызов (или любая другая операция XLA) возвращает несколько независимых массивов.
Но менее очевидно, что вывод в виде кортежа также является способом предоставления вашему пользовательскому вызову временной памяти. Да, выходной параметр может представлять собой временный буфер. Рассмотрим пример: выходной буфер обладает свойством, позволяющим операции записывать в него данные и считывать их после записи. Это именно то, что вам нужно от временного буфера.
В приведенном выше примере предположим, что мы хотим использовать F32[1024] в качестве временного буфера. Тогда мы бы записали HLO точно так же, как и выше, и просто никогда бы не читали кортеж с индексом 1 из выходных данных пользовательского вызова.