การเรียกที่กำหนดเองของ XLA

เอกสารนี้อธิบายวิธีเขียนและใช้การเรียกที่กำหนดเองของ XLA โดยใช้ไลบรารี XLA FFI การเรียกที่กำหนดเองเป็นกลไกในการอธิบาย "การดำเนินการ" ภายนอกในโมดูล HLO ให้คอมไพเลอร์ XLA (ในเวลาคอมไพล์) และ XLA FFI เป็นกลไกในการลงทะเบียนการใช้งานการดำเนินการดังกล่าวกับ XLA (ในเวลารัน) FFI ย่อมาจาก "foreign function interface" และเป็นชุด API ของ C ที่กำหนดอินเทอร์เฟซไบนารี (ABI) สำหรับ XLA เพื่อเรียกใช้โค้ดภายนอกที่เขียนด้วยภาษาโปรแกรมอื่นๆ XLA มีการผูกเฉพาะส่วนหัวสำหรับ XLA FFI ที่เขียนด้วย C++ ซึ่งซ่อนรายละเอียดระดับต่ำทั้งหมดของ API ของ C ที่อยู่เบื้องหลังจากผู้ใช้ปลายทาง

การเรียกที่กำหนดเองของ JAX + XLA

ดูเอกสารประกอบของ JAX สำหรับ ตัวอย่างแบบครบวงจรของการผสานรวมการเรียกที่กำหนดเองและ XLA FFI กับ JAX

การผูก XLA FFI

การผูก XLA FFI เป็นข้อกำหนดเวลาคอมไพล์ของลายเซ็นการเรียกที่กำหนดเอง ได้แก่ อาร์กิวเมนต์ แอตทริบิวต์ และประเภทของการเรียกที่กำหนดเอง รวมถึงพารามิเตอร์เพิ่มเติมที่ส่งผ่านบริบทการดำเนินการ (เช่น สตรีม GPU สำหรับแบ็กเอนด์ GPU) การผูก XLA FFI สามารถผูกกับ C++ ที่เรียกได้ (ตัวชี้ฟังก์ชัน, แลมบ์ดา ฯลฯ) ที่มีลายเซ็น operator() ที่เข้ากันได้ ตัวแฮนเดิลที่สร้างขึ้นจะถอดรหัสเฟรมการเรียก XLA FFI (กำหนดโดย C API ที่เสถียร) ตรวจสอบประเภทพารามิเตอร์ทั้งหมด และส่งต่อผลลัพธ์ที่ถอดรหัสแล้วไปยังการเรียกกลับที่ผู้ใช้กำหนด

การผูก XLA FFI อาศัยการเขียนโปรแกรมเมตาแบบเทมเพลตอย่างมากเพื่อให้สามารถคอมไพล์ตัวแฮนเดิลที่สร้างขึ้นเป็นรหัสเครื่องที่มีประสิทธิภาพสูงสุด ค่าใช้จ่ายในการรันไทม์จะอยู่ที่ประมาณ 2-3 นาโนวินาทีสำหรับพารามิเตอร์การเรียกที่กำหนดเองแต่ละรายการ

จุดการปรับแต่ง 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 ให้เป็นประเภทบัฟเฟอร์ที่เฉพาะเจาะจง การตรวจสอบขนาดมิติข้อมูล และการตรวจสอบความสัมพันธ์ระหว่างรูปร่างบัฟเฟอร์หลายรายการ

รูปแบบบัฟเฟอร์ไม่สามารถเปลี่ยนแปลงได้ และสามารถระบุ dtype และอันดับได้โดยตรงหรือใช้ตัวแก้ไขที่เกี่ยวข้อง

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 จะบันทึกขนาดมิติข้อมูล 2 รายการ และการจับคู่ 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 ไม่จำกัด WithDim(index, ...) จะยอมรับอันดับใดก็ได้ที่มากกว่า index หากไม่มี WithRank ที่ชัดเจน

การบันทึกมิติข้อมูลยังแสดงความสัมพันธ์ระหว่างตำแหน่งได้ด้วย

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));
}

ตัวแก้ไขทั้ง 2 รายการจะคัดลอกข้อมูลเมตาของบัฟเฟอร์อ้างอิงเมื่อสร้างรูปแบบ โดยรูปแบบจะไม่เก็บการอ้างอิงไปยังบัฟเฟอร์ เนื่องจากข้อจำกัดเหล่านี้เป็นข้อจำกัดรันไทม์ จึงไม่สามารถปรับแต่ง AnyBuffer ให้เป็นประเภทการแสดงผลที่เฉพาะเจาะจงได้ ข้อจำกัดเหล่านี้มีประโยชน์หลักๆ กับ Verify หรือกับ Match เมื่อบัฟเฟอร์อินพุตมีประเภทที่เฉพาะเจาะจงอยู่แล้ว

ใช้ Verify เมื่อบัฟเฟอร์ควรคงประเภทที่มีอยู่ไว้ หรือเมื่อรูปแบบยอมรับประเภทข้อมูลหรืออันดับมากกว่า 1 รายการ ตัวอย่างเช่น บัฟเฟอร์ดัชนียอมรับ 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 ที่ส่งผ่านเป็น backend_config ของ custom_call เป็นอาร์กิวเมนต์ตัวแฮนเดิล 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>

ในตัวอย่างนี้ การเรียกที่กำหนดเองมีอาร์กิวเมนต์บัฟเฟอร์เดียวและแอตทริบิวต์ 2 รายการ และ 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 ที่เป็นจำนวนเต็มเป็น Enum ที่ผู้ใช้กำหนดโดยอัตโนมัติ คลาส Enum ต้องมีประเภทจำนวนเต็มพื้นฐานเดียวกัน และต้องลงทะเบียนการถอดรหัสกับ 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();
});

แอตทริบิวต์ Struct ที่ผู้ใช้กำหนด

XLA FFI สามารถถอดรหัสแอตทริบิวต์พจนานุกรมเป็น Struct ที่ผู้ใช้กำหนด

%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 และแทนที่จะเข้าถึงฟิลด์พจนานุกรมตามชื่อ ระบบจะถอดรหัสแอตทริบิวต์นี้เป็น Struct ของ 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();
});

สร้างการเรียกที่กำหนดเองใน CPU

คุณสามารถสร้างคำสั่ง HLO ที่แสดงการเรียกที่กำหนดเองผ่าน Client API ของ XLA ได้ ตัวอย่างเช่น โค้ดต่อไปนี้ใช้การเรียกที่กำหนดเองเพื่อคำนวณ A[i] = B[i % 128]+ C[i] ใน CPU (แน่นอนว่าคุณสามารถทำได้และควรทำ! ด้วย 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

การลงทะเบียนการเรียกที่กำหนดเองของ 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 ฟังก์ชัน CPU do_custom_call มีหน้าที่รับผิดชอบในการจัดคิวงานใน GPU ในที่นี้ ฟังก์ชันจะเปิดใช้เคอร์เนล CUDA แต่ก็อาจทำอย่างอื่นได้ด้วย เช่น เรียก cuBLAS

อาร์กิวเมนต์และผลลัพธ์จะอยู่ในโฮสต์ด้วย และสมาชิกข้อมูลจะมีตัวชี้ไปยังหน่วยความจำของอุปกรณ์ (เช่น GPU) บัฟเฟอร์ที่ส่งผ่านไปยังตัวแฮนเดิลการเรียกที่กำหนดเองจะมีรูปร่างเหมือนกับบัฟเฟอร์อุปกรณ์พื้นฐาน ดังนั้นการเรียกที่กำหนดเองจึงสามารถคำนวณพารามิเตอร์การเปิดใช้เคอร์เนลจากบัฟเฟอร์เหล่านี้ได้

การส่งผ่านทูเพิลไปยังการเรียกที่กำหนดเอง

ลองพิจารณาการเรียกที่กำหนดเองต่อไปนี้

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, ...);

ทั้งใน CPU และ GPU ระบบจะแสดงทูเพิลในหน่วยความจำเป็นอาร์เรย์ของตัวชี้ เมื่อ XLA เรียกการเรียกที่กำหนดเองด้วยอาร์กิวเมนต์หรือผลลัพธ์ของทูเพิล ระบบจะทำให้ทูเพิลแบนลงและส่งผ่านเป็นอาร์กิวเมนต์หรือผลลัพธ์ของบัฟเฟอร์ปกติ

เอาต์พุตทูเพิลเป็นบัฟเฟอร์ชั่วคราว

อินพุตทูเพิลไปยังการเรียกที่กำหนดเองเป็นสิ่งอำนวยความสะดวก แต่ไม่ใช่สิ่งที่จำเป็น หากเราไม่รองรับอินพุตทูเพิลไปยังการเรียกที่กำหนดเอง คุณสามารถแกะทูเพิลได้ทุกเมื่อโดยใช้ get-tuple-element ก่อนที่จะส่งทูเพิลไปยังการเรียกที่กำหนดเอง

ในทางกลับกัน เอาต์พุตทูเพิลช่วยให้คุณทำสิ่งต่างๆ ที่ทำไม่ได้ในกรณีอื่นๆ

เหตุผลที่ชัดเจนในการมีเอาต์พุตทูเพิลคือเอาต์พุตทูเพิลเป็นวิธีที่การเรียกที่กำหนดเอง (หรือการดำเนินการ XLA อื่นๆ) แสดงผลอาร์เรย์อิสระหลายรายการ

แต่สิ่งที่ชัดเจนน้อยกว่าคือเอาต์พุตทูเพิลยังเป็นวิธีในการให้หน่วยความจำชั่วคราวแก่การเรียกที่กำหนดเองด้วย ใช่ เอาต์พุตสามารถแสดงบัฟเฟอร์ชั่วคราวได้ ลองพิจารณาว่าบัฟเฟอร์เอาต์พุตมีพร็อพเพอร์ตี้ที่การดำเนินการสามารถเขียนลงในบัฟเฟอร์และอ่านจากบัฟเฟอร์ได้หลังจากเขียนลงไปแล้ว ซึ่งเป็นสิ่งที่คุณต้องการจากบัฟเฟอร์ชั่วคราว

ในตัวอย่างด้านบน สมมติว่าเราต้องการใช้ F32[1024] เป็นบัฟเฟอร์ชั่วคราว จากนั้นเราจะเขียน HLO เหมือนกับด้านบน และจะไม่อ่านดัชนีทูเพิล 1 ของเอาต์พุตการเรียกที่กำหนดเอง