Ce document explique comment écrire et utiliser des appels personnalisés XLA à l'aide de la bibliothèque XLA FFI. Un appel personnalisé est un mécanisme permettant de décrire une "opération" externe dans le module HLO au compilateur XLA (au temps de compilation). XLA FFI est un mécanisme permettant d'enregistrer l'implémentation de ces opérations avec XLA (au moment de l'exécution). FFI signifie "foreign function interface" (interface de fonction externe). Il s'agit d'un ensemble d'API C qui définissent une interface binaire (ABI) pour qu'XLA puisse appeler du code externe écrit dans d'autres langages de programmation. XLA fournit des liaisons d'en-tête uniquement pour XLA FFI écrit en C++, ce qui masque tous les détails de bas niveau des API C sous-jacentes à l'utilisateur final.
Appels personnalisés JAX + XLA
Consultez la documentation JAX pour obtenir des exemples de bout en bout d'intégration d'appels personnalisés et de XLA FFI avec JAX.
Liaison XLA FFI
La liaison XLA FFI est une spécification au moment de la compilation de la signature d'appel personnalisé : arguments d'appel personnalisé, attributs et leurs types, et paramètres supplémentaires transmis via le contexte d'exécution (c'est-à-dire le flux GPU pour le backend GPU). La liaison XLA FFI peut être liée à n'importe quel élément C++ appelable (pointeur de fonction, lambda, etc.) avec une signature operator() compatible. Le gestionnaire construit décode le frame d'appel XLA FFI (défini par l'API C stable), vérifie le type de tous les paramètres et transmet les résultats décodés au rappel défini par l'utilisateur.
La liaison XLA FFI repose fortement sur la métaprogrammation de modèle pour pouvoir compiler le gestionnaire construit dans le code machine le plus efficace. Les frais généraux d'exécution sont de l'ordre de quelques nanosecondes pour chaque paramètre d'appel personnalisé.
Les points de personnalisation XLA FFI sont implémentés en tant que spécialisations de modèle, et les utilisateurs peuvent définir comment décoder leurs types personnalisés. Il est donc possible de définir un décodage personnalisé pour les types enum class définis par l'utilisateur.
Renvoi d'erreurs à partir d'appels personnalisés
Les implémentations d'appel personnalisé doivent renvoyer la valeur xla::ffi::Error pour signaler la réussite ou l'échec à l'exécution XLA. Elle est semblable à absl::Status et comporte le même ensemble de codes d'erreur. Nous n'utilisons pas absl::Status, car elle ne dispose pas d'une ABI stable et il serait dangereux de la transmettre entre la bibliothèque d'appel personnalisé chargée dynamiquement et XLA lui-même.
// 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(); });
Arguments et résultats de mémoire tampon
XLA utilise le style de transmission de destination pour les résultats : les appels personnalisés (ou toute autre opération XLA) n'allouent pas de mémoire pour les résultats, mais écrivent plutôt dans les destinations transmises par l'exécution XLA. XLA utilise l'attribution de mémoire tampon statique et alloue des mémoires tampons pour toutes les valeurs en fonction de leurs plages actives au temps de compilation.
Les résultats transmis aux gestionnaires FFI sont encapsulés dans un Result<T> modèle, qui
présente une sémantique semblable à celle d'un pointeur : operator-> donne accès au paramètre sous-jacent.
Les arguments et les résultats AnyBuffer donnent accès aux paramètres de mémoire tampon d'appel personnalisé de n'importe quel type de données. Cela est utile lorsque l'appel personnalisé dispose d'une implémentation générique qui fonctionne pour plusieurs types de données, et que l'implémentation d'appel personnalisé effectue une distribution au moment de l'exécution en fonction du type de données. AnyBuffer donne accès au type de données de la mémoire tampon, aux dimensions et à un pointeur vers la mémoire tampon elle-même.
%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();
});
Arguments et résultats de mémoire tampon contraints
Buffer permet d'ajouter des contraintes sur le type de données de la mémoire tampon et le nombre de dimensions. Elles seront automatiquement vérifiées par le gestionnaire et renverront une erreur à l'exécution XLA si les arguments d'exécution ne correspondent pas à la signature du gestionnaire 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();
});
Mise en correspondance et vérification des mémoires tampons
Les modèles de mémoire tampon peuvent exprimer des contraintes plus spécifiques que les types d'arguments et de résultats dans une liaison FFI. Ils sont utiles pour affiner un AnyBuffer en un type de mémoire tampon concret, vérifier les tailles de dimension et vérifier les relations entre plusieurs formes de mémoire tampon.
Les modèles sont immuables et peuvent spécifier le dtype et le rang directement ou avec les modificateurs correspondants :
namespace m = ::xla::ffi::match;
m::Buffer<F32, 2>();
m::Buffer().WithDType<F32>().WithRank<2>();
La fonction Match vérifie un modèle et renvoie son type de mémoire tampon le plus contraint. Pour affiner un AnyBuffer, le modèle doit spécifier exactement un type de données et un rang. La mise en correspondance d'une mémoire tampon déjà typée vérifie toutes les contraintes supplémentaires et conserve son type.
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();
});
Dans cet exemple, la mise en correspondance de input capture ses deux tailles de dimension, et la mise en correspondance de output vérifie qu'elle a la même forme. Les captures ne sont validées que lorsque le modèle complet correspond.
Un modèle de dimension peut être non contraint, fixe ou capturé. Les valeurs fixes et les pointeurs de capture sont implicitement convertis en modèles de dimension :
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 décrit une forme complète et définit son rang à partir du nombre d'arguments. WithDim<I> ou WithDim(index, ...) contraint les dimensions
de manière positionnelle sans fixer le rang complet. Une contrainte positionnelle nécessite l'existence de la dimension référencée, mais laisse toutes les autres dimensions non contraintes :
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 nécessite une mémoire tampon F32 de rang 4 avec des dimensions 0 et 3 égales à 1 et 128, respectivement. Les dimensions 1 et 2 ne sont pas limitées. Sans un
explicite WithRank, WithDim(index, ...) accepte tout rang supérieur à
index.
Les captures de dimension peuvent également exprimer des relations entre les positions :
auto matrix = m::Buffer<F32>().WithDims(&rows, 128);
int64_t n;
ErrorOr<BufferR2<F32>> square =
Match("matrix", buffer,
m::Buffer<F32>().WithDims(&n, &n));
Lorsque le même pointeur de capture apparaît plusieurs fois, toutes les dimensions correspondantes doivent avoir la même taille. L'appel Match ci-dessus n'accepte donc que les matrices carrées. Les captures de dimension ne sont écrites qu'une fois le modèle complet réussi et restent inchangées en cas d'échec.
Les relations entre les mémoires tampons entières ne nécessitent pas la capture de chaque dimension.
WithShapeOf correspond à la forme d'exécution complète d'une autre mémoire tampon tout en autorisant un dtype différent. Like nécessite également le même dtype :
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));
}
Les deux modificateurs copient les métadonnées de la mémoire tampon de référence lors de la création du modèle. Le modèle ne conserve pas de référence à la mémoire tampon. Comme il s'agit de contraintes d'exécution, elles n'affinent pas un AnyBuffer en un type renvoyé concret. Elles sont principalement utiles avec Verify ou avec Match lorsque la mémoire tampon d'entrée a déjà un type concret.
Utilisez Verify lorsque la mémoire tampon doit conserver son type existant ou lorsqu'un modèle accepte plusieurs types de données ou rangs. Par exemple, une mémoire tampon d'index peut accepter S32 ou S64 et distribuer son type de données après vérification :
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();
}
Comme ce modèle autorise plusieurs types de mémoire tampon concrets, il ne peut pas affiner un AnyBuffer avec Match : il n'existe pas de type renvoyé unique. Il peut toujours être transmis à Match avec une mémoire tampon déjà typée, dont le type renvoyé est déjà connu. Une mémoire tampon déjà typée peut également être vérifiée. Cela est généralement utile pour vérifier les relations de forme :
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));
}
Le nom transmis en tant que premier argument à Match et Verify est obligatoire et est inclus dans les erreurs, ainsi qu'une description de la contrainte ayant échoué.
L'API FFI externe dans xla/ffi/api/ffi.h renvoie Error à partir de Verify et
ErrorOr<T> à partir de Match, comme dans les exemples ci-dessus. L'API FFI interne dans
xla/ffi/ffi.h fournit la même interface de mise en correspondance à l'aide de absl::Status et
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())));
Arguments et résultats variadiques
Si le nombre d'arguments et de résultats peut être différent dans différentes instances d'un appel personnalisé, ils peuvent être décodés au moment de l'exécution à l'aide de RemainingArgs et 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();
});
Les arguments et les résultats variadiques peuvent être déclarés après les arguments et les résultats réguliers. Toutefois, la liaison d'arguments et de résultats réguliers après un argument variadique est illégale.
auto handler =
Ffi::Bind()
.Arg<AnyBuffer>()
.RemainingArgs()
.Ret<AnyBuffer>()
.RemainingRets()
.To([](AnyBuffer arg, RemainingArgs args, AnyBuffer ret,
RemainingRets results) -> Error { return Error::Success(); });
Attributs
XLA FFI est compatible avec le décodage automatique de mlir::DictionaryAttr transmis en tant que custom_call backend_config dans les arguments du gestionnaire 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>
Dans cet exemple, l'appel personnalisé comporte un seul argument de mémoire tampon et deux attributs. XLA FFI peut les décoder automatiquement et les transmettre à l'élément appelable défini par l'utilisateur.
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();
});
Attributs d'énumération définis par l'utilisateur
XLA FFI peut décoder automatiquement les attributs MLIR intégraux en énumérations définies par l'utilisateur. La classe d'énumération doit avoir le même type intégral sous-jacent, et le décodage doit être explicitement enregistré auprès de 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(); });
Liaison de tous les attributs d'appel personnalisé
Il est possible d'accéder à tous les attributs d'appel personnalisé en tant que dictionnaire et de ne décoder que les attributs nécessaires au moment de l'exécution.
auto handler = Ffi::Bind().Attrs().To([](Dictionary attrs) -> Error {
ErrorOr<int32_t> i32 = attrs.get<int32_t>("i32");
return Error::Success();
});
Attributs de structure définis par l'utilisateur
XLA FFI peut décoder les attributs de dictionnaire en structures définies par l'utilisateur.
%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>
Dans l'exemple ci-dessus, range est un attribut mlir::DictionaryAttr. Au lieu d'accéder aux champs de dictionnaire par nom, il peut être décodé automatiquement en tant que structure C++. Le décodage doit être explicitement enregistré avec une macro XLA_FFI_REGISTER_STRUCT_ATTR_DECODING (en arrière-plan, il définit une spécialisation de modèle dans l'espace de noms ::xla::ffi. La macro doit donc être ajoutée à l'espace de noms global).
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();
});
Les attributs personnalisés peuvent être chargés à partir d'un dictionnaire, comme n'importe quel autre attribut. Dans l'exemple ci-dessous, tous les attributs d'appel personnalisé sont décodés en tant que Dictionary, et un range est accessible par nom.
auto handler = Ffi::Bind().Attrs().To([](Dictionary attrs) -> Error {
ErrorOr<Range> range = attrs.get<Range>("range");
return Error::Success();
});
Créer un appel personnalisé sur le processeur
Vous pouvez créer une instruction HLO qui représente un appel personnalisé via l'API cliente de XLA. Par exemple, le code suivant utilise un appel personnalisé pour calculer A[i] = B[i %
128]+ C[i] sur le processeur. (Bien sûr, vous pouvez et devez le faire avec un HLO normal.)
#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);
Créer un appel personnalisé sur le GPU
L'enregistrement d'appel personnalisé GPU avec XLA FFI est presque identique. La seule différence est que pour le GPU, vous devez demander un flux de plate-forme sous-jacent (flux CUDA ou ROCM) pour pouvoir lancer le noyau sur l'appareil. Voici un exemple CUDA qui effectue le même calcul (A[i] = B[i % 128] + C[i]) que le code du processeur ci-dessus.
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);
Notez tout d'abord que la fonction d'appel personnalisé GPU est toujours une fonction exécutée sur le processeur. La fonction de processeur do_custom_call est chargée de mettre en file d'attente le travail sur le GPU. Ici, elle lance un noyau CUDA, mais elle peut également effectuer une autre action, comme appeler cuBLAS.
Les arguments et les résultats se trouvent également sur l'hôte, et le membre de données contient un pointeur vers la mémoire de l'appareil (c'est-à-dire le GPU). Les mémoires tampons transmises au gestionnaire d'appel personnalisé ont la forme des mémoires tampons de l'appareil sous-jacent. L'appel personnalisé peut donc calculer les paramètres de lancement du noyau à partir de celles-ci.
Transmettre des tuples à des appels personnalisés
Prenons l'appel personnalisé suivant.
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, ...);
Sur le processeur et le GPU, un tuple est représenté en mémoire sous la forme d'un tableau de pointeurs. Lorsque XLA appelle des appels personnalisés avec des arguments ou des résultats de tuple, il les aplatit et les transmet en tant qu'arguments ou résultats de mémoire tampon réguliers.
Sorties de tuple en tant que mémoires tampons temporaires
Les entrées de tuple dans les appels personnalisés sont pratiques, mais elles ne sont pas strictement nécessaires. Si nous n'étions pas compatibles avec les entrées de tuple dans les appels personnalisés, vous pourriez toujours décompresser les tuples à l'aide de get-tuple-element avant de les transmettre à l'appel personnalisé.
En revanche, les sorties de tuple vous permettent d'effectuer des actions que vous ne pourriez pas faire autrement.
La raison évidente d'avoir des sorties de tuple est qu'elles permettent à un appel personnalisé (ou à toute autre opération XLA) de renvoyer plusieurs tableaux indépendants.
Mais moins évidemment, une sortie de tuple est également un moyen de donner à votre appel personnalisé une mémoire temporaire. Oui, une sortie peut représenter une mémoire tampon temporaire. Considérez qu'une mémoire tampon de sortie a la propriété que l'opération peut y écrire et qu'elle peut y lire après avoir été écrite. C'est exactement ce que vous attendez d'une mémoire tampon temporaire.
Dans l'exemple ci-dessus, supposons que nous voulions utiliser F32[1024] comme mémoire tampon temporaire.
Nous écririons alors le HLO comme ci-dessus, et nous ne lirions jamais l'index de tuple 1 de la sortie de l'appel personnalisé.