এই ডকুমেন্টটিতে XLA FFI লাইব্রেরি ব্যবহার করে কীভাবে XLA কাস্টম কল লিখতে ও ব্যবহার করতে হয়, তা বর্ণনা করা হয়েছে। কাস্টম কল হলো HLO মডিউলের কোনো বাহ্যিক "অপারেশন" XLA কম্পাইলারের কাছে (কম্পাইল করার সময়) বর্ণনা করার একটি পদ্ধতি, এবং XLA FFI হলো এই ধরনের অপারেশনের ইমপ্লিমেন্টেশন XLA-এর সাথে (রান করার সময়) রেজিস্টার করার একটি পদ্ধতি। FFI-এর পূর্ণরূপ হলো "ফরেন ফাংশন ইন্টারফেস" এবং এটি C API-এর একটি সেট যা অন্য প্রোগ্রামিং ভাষায় লেখা বাহ্যিক কোড কল করার জন্য XLA-এর একটি বাইনারি ইন্টারফেস (ABI) নির্ধারণ করে। XLA, C++-এ লেখা XLA FFI-এর জন্য শুধুমাত্র হেডার বাইন্ডিং সরবরাহ করে, যা ব্যবহারকারীর কাছ থেকে অন্তর্নিহিত C API-এর সমস্ত নিম্ন-স্তরের বিবরণ গোপন রাখে।
JAX + XLA কাস্টম কল
JAX-এর সাথে কাস্টম কল এবং XLA FFI ইন্টিগ্রেট করার সম্পূর্ণ উদাহরণের জন্য JAX ডকুমেন্টেশন দেখুন।
XLA FFI বাইন্ডিং
XLA FFI বাইন্ডিং হলো কাস্টম কল সিগনেচারের একটি কম্পাইল-টাইম স্পেসিফিকেশন: এতে থাকে কাস্টম কল আর্গুমেন্ট, অ্যাট্রিবিউট ও তাদের টাইপ, এবং এক্সিকিউশন কনটেক্সটের (যেমন, GPU ব্যাকএন্ডের জন্য জিপিইউ স্ট্রিম) মাধ্যমে পাঠানো অতিরিক্ত প্যারামিটার। XLA FFI বাইন্ডিংকে সামঞ্জস্যপূর্ণ operator() সিগনেচারযুক্ত যেকোনো C++ কলযোগ্য উপাদানের (ফাংশন পয়েন্টার, ল্যাম্বডা, ইত্যাদি) সাথে বাইন্ড করা যায়। নির্মিত হ্যান্ডলারটি XLA FFI কল ফ্রেম (যা স্টেবল C API দ্বারা সংজ্ঞায়িত) ডিকোড করে, সমস্ত প্যারামিটারের টাইপ চেক করে এবং ডিকোড করা ফলাফল ব্যবহারকারী-সংজ্ঞায়িত কলব্যাকে ফরোয়ার্ড করে।
নির্মিত হ্যান্ডলারকে সবচেয়ে কার্যকর মেশিন কোডে কম্পাইল করার জন্য XLA FFI বাইন্ডিং টেমপ্লেট মেটাপ্রোগ্রামিংয়ের ওপর ব্যাপকভাবে নির্ভর করে। প্রতিটি কাস্টম কল প্যারামিটারের জন্য রান টাইম ওভারহেড কয়েক ন্যানোসেকেন্ডের মতো হয়ে থাকে।
XLA FFI কাস্টমাইজেশন পয়েন্টগুলো টেমপ্লেট স্পেশালাইজেশন হিসেবে প্রয়োগ করা হয়, এবং ব্যবহারকারীরা তাদের কাস্টম টাইপগুলো কীভাবে ডিকোড করবে তা নির্ধারণ করতে পারে; অর্থাৎ, ব্যবহারকারী-সংজ্ঞায়িত enum class টাইপগুলোর জন্য কাস্টম ডিকোডিং নির্ধারণ করা সম্ভব।
কাস্টম কল থেকে ত্রুটি ফেরত আসা
কাস্টম কল ইমপ্লিমেন্টেশনগুলোকে অবশ্যই XLA রানটাইমকে সাফল্য বা ত্রুটির সংকেত দিতে xla::ffi::Error ভ্যালুটি রিটার্ন করতে হবে। এটি 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 , এবং হ্যান্ডলার দ্বারা সেগুলি স্বয়ংক্রিয়ভাবে যাচাই করা হবে ও রানটাইম আর্গুমেন্টগুলো FFI হ্যান্ডলার সিগনেচারের সাথে না মিললে XLA রানটাইমে একটি ত্রুটি ফেরত পাঠানো হবে।
// 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 একটি নির্দিষ্ট বাফার টাইপে পরিমার্জন করতে, ডাইমেনশনের আকার পরীক্ষা করতে এবং একাধিক বাফার শেপের মধ্যে সম্পর্ক যাচাই করতে উপযোগী।
প্যাটার্নগুলি অপরিবর্তনীয় এবং সরাসরি অথবা সংশ্লিষ্ট মডিফায়ারগুলির সাথে dtype ও rank নির্দিষ্ট করতে পারে:
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 বাফার প্রয়োজন, যার ডাইমেনশন ০ এবং ৩ যথাক্রমে ১ এবং ১২৮-এর সমান হবে; ডাইমেনশন ১ এবং ২-এর উপর কোনো বিধিনিষেধ নেই। সুস্পষ্ট 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 অন্য একটি বাফারের সম্পূর্ণ রানটাইম শেপের সাথে মেলে এবং একই সাথে একটি ভিন্ন dtype ব্যবহারের সুযোগ দেয়। Like ক্ষেত্রেও একই 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));
}
প্যাটার্নটি তৈরি করার সময় উভয় মডিফায়ারই রেফারেন্স বাফারের মেটাডেটা কপি করে; প্যাটার্নটি বাফারটির কোনো রেফারেন্স ধরে রাখে না। যেহেতু এগুলো রানটাইম সীমাবদ্ধতা, তাই এগুলো কোনো 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();
}
যেহেতু এই প্যাটার্নটি একাধিক কংক্রিট বাফার টাইপ অনুমোদন করে, তাই এটি Match ব্যবহার করে একটি AnyBuffer পরিমার্জন করতে পারে না: এর কোনো অনন্য রিটার্ন টাইপ নেই। তবে, এটিকে আগে থেকে টাইপ করা একটি বাফারের সাথে 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 এর প্রথম আর্গুমেন্ট হিসেবে প্রদত্ত নামটি আবশ্যক এবং ব্যর্থ হওয়া কনস্ট্রেইন্টের বিবরণের সাথে এটিও ত্রুটির তালিকায় অন্তর্ভুক্ত থাকে।
উপরের উদাহরণগুলোর মতো, xla/ffi/api/ffi.h এ থাকা এক্সটার্নাল FFI API, Verify থেকে Error এবং Match থেকে ErrorOr<T> রিটার্ন করে। xla/ffi/ffi.h এ থাকা ইন্টারনাল FFI API, 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, FFI হ্যান্ডলার আর্গুমেন্টে custom_call backend_config হিসেবে পাস করা mlir::DictionaryAttr এর স্বয়ংক্রিয় ডিকোডিং সমর্থন করে।
%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();
});
ব্যবহারকারী-সংজ্ঞায়িত Enum অ্যাট্রিবিউট
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++ struct হিসেবে ডিকোড করা যায়। এই ডিকোডিং প্রক্রিয়াটি একটি 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();
});
সিপিইউতে একটি কাস্টম কল তৈরি করুন
আপনি XLA-এর ক্লায়েন্ট API ব্যবহার করে একটি কাস্টম কলকে প্রতিনিধিত্বকারী একটি HLO ইন্সট্রাকশন তৈরি করতে পারেন। উদাহরণস্বরূপ, নিম্নলিখিত কোডটি CPU-তে 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 কাস্টম কল রেজিস্ট্রেশন প্রায় একই রকম, একমাত্র পার্থক্য হলো GPU-এর ক্ষেত্রে ডিভাইসে কার্নেল চালু করার জন্য একটি অন্তর্নিহিত প্ল্যাটফর্ম স্ট্রিমের (CUDA বা ROCM স্ট্রিম) জন্য অনুরোধ করতে হয়। এখানে একটি CUDA উদাহরণ দেওয়া হলো যা উপরের CPU কোডের মতোই একই গণনা ( A[i] = B[i % 128] + C[i] ) করে।
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 লিখব, এবং আমরা কাস্টম কলের আউটপুটের টাপল ইনডেক্স ১ কখনোই পড়ব না।