0 | {--
  1 | Copyright (C) 2024  Joel Berkeley
  2 |
  3 | This program is free software: you can redistribute it and/or modify
  4 | it under the terms of the GNU Affero General Public License as published
  5 | by the Free Software Foundation, either version 3 of the License, or
  6 | (at your option) any later version.
  7 |
  8 | This program is distributed in the hope that it will be useful,
  9 | but WITHOUT ANY WARRANTY; without even the implied warranty of
 10 | MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
 11 | GNU Affero General Public License for more details.
 12 |
 13 | You should have received a copy of the GNU Affero General Public License
 14 | along with this program.  If not, see <https://www.gnu.org/licenses/>.
 15 | --}
 16 | ||| For internal spidr use, and use by plugin developers.
 17 | |||
 18 | ||| The Idris API for PJRT.
 19 | module Compiler.Xla.PJRT.C.PjrtCApi
 20 |
 21 | import public Control.Monad.Either
 22 | import Derive.Prelude
 23 | import Language.Reflection
 24 |
 25 | import Compiler.FFI
 26 | import Compiler.String
 27 | import Compiler.Array
 28 | import Types
 29 | import Util
 30 |
 31 | %language ElabReflection
 32 |
 33 | ffi : String -> String
 34 | ffi = libxla "c/xla/pjrt/c/pjrt_c_api.h"
 35 |
 36 | ||| For use by plugin developers.
 37 | |||
 38 | ||| A minimal wrapper round a C `PJRT_Api` struct pointer. The memory should be owned by the
 39 | ||| code producing the pointer.
 40 | public export
 41 | data PjrtApi = MkPjrtApi AnyPtr
 42 |
 43 | ||| The cause of a `PjrtError`.
 44 | public export
 45 | data PjrtErrorCode =
 46 |     PJRT_Error_Code_CANCELLED
 47 |   | PJRT_Error_Code_UNKNOWN
 48 |   | PJRT_Error_Code_INVALID_ARGUMENT
 49 |   | PJRT_Error_Code_DEADLINE_EXCEEDED
 50 |   | PJRT_Error_Code_NOT_FOUND
 51 |   | PJRT_Error_Code_ALREADY_EXISTS
 52 |   | PJRT_Error_Code_PERMISSION_DENIED
 53 |   | PJRT_Error_Code_RESOURCE_EXHAUSTED
 54 |   | PJRT_Error_Code_FAILED_PRECONDITION
 55 |   | PJRT_Error_Code_ABORTED
 56 |   | PJRT_Error_Code_OUT_OF_RANGE
 57 |   | PJRT_Error_Code_UNIMPLEMENTED
 58 |   | PJRT_Error_Code_INTERNAL
 59 |   | PJRT_Error_Code_UNAVAILABLE
 60 |   | PJRT_Error_Code_DATA_LOSS
 61 |   | PJRT_Error_Code_UNAUTHENTICATED
 62 |
 63 | %runElab derive "PjrtErrorCode" [Show]
 64 |
 65 | ||| Indicates an error in the PJRT C layer, either due to internal errors or user error.
 66 | public export
 67 | record PjrtError where
 68 |   constructor MkPjrtError
 69 |
 70 |   ||| The error message.
 71 |   message : String
 72 |
 73 |   ||| The error cause code, if one exists.
 74 |   code : Maybe PjrtErrorCode
 75 |
 76 | export
 77 | Show PjrtError where
 78 |   show e =
 79 |     let code = case e.code of
 80 |           Nothing => "unknown"
 81 |           Just c => show c
 82 |      in "PjrtError (error code \{code})\n\{e.message}"
 83 |
 84 | %foreign (ffi "PJRT_Error_Destroy_Args_delete")
 85 | prim__deletePjrtErrorDestroyArgs : AnyPtr -> PrimIO ()
 86 |
 87 | %foreign (ffi "PJRT_Error_Destroy_Args_new")
 88 | prim__mkPjrtErrorDestroyArgs : AnyPtr -> PrimIO AnyPtr
 89 |
 90 | %foreign (ffi "pjrt_error_destroy")
 91 | prim__pjrtErrorDestroy : AnyPtr -> AnyPtr -> PrimIO ()
 92 |
 93 | destroyPjrtError : HasIO io => AnyPtr -> AnyPtr -> io ()
 94 | destroyPjrtError api err = do
 95 |   args <- primIO $ prim__mkPjrtErrorDestroyArgs err
 96 |   primIO $ prim__pjrtErrorDestroy api args
 97 |   primIO $ prim__deletePjrtErrorDestroyArgs args
 98 |
 99 | %foreign (ffi "PJRT_Error_Message_Args_delete")
100 | prim__deletePjrtErrorMessageArgs : AnyPtr -> PrimIO ()
101 |
102 | %foreign (ffi "PJRT_Error_Message_Args_new")
103 | prim__mkPjrtErrorMessageArgs : AnyPtr -> PrimIO AnyPtr
104 |
105 | %foreign (ffi "PJRT_Error_Message_Args_message")
106 | prim__pjrtErrorMessageArgsMessage : AnyPtr -> PrimIO String
107 |
108 | %foreign (ffi "pjrt_error_message")
109 | prim__pjrtErrorMessage : AnyPtr -> AnyPtr -> PrimIO ()
110 |
111 | pjrtErrorMessage : HasIO io => AnyPtr -> AnyPtr -> io String
112 | pjrtErrorMessage api err = do
113 |   args <- primIO $ prim__mkPjrtErrorMessageArgs err
114 |   primIO $ prim__pjrtErrorMessage api args
115 |   msg <- primIO $ prim__pjrtErrorMessageArgsMessage args
116 |   primIO $ prim__deletePjrtErrorMessageArgs args
117 |   pure msg
118 |
119 | %foreign (ffi "PJRT_Error_GetCode_Args_delete")
120 | prim__deletePjrtErrorGetCodeArgs : AnyPtr -> PrimIO ()
121 |
122 | %foreign (ffi "PJRT_Error_GetCode_Args_new")
123 | prim__mkPjrtErrorGetCodeArgs : AnyPtr -> PrimIO AnyPtr
124 |
125 | %foreign (ffi "PJRT_Error_GetCode_Args_code")
126 | prim__pjrtErrorGetCodeArgsCode : AnyPtr -> Int
127 |
128 | %foreign (ffi "pjrt_error_getcode")
129 | prim__pjrtErrorGetCode : AnyPtr -> AnyPtr -> PrimIO AnyPtr
130 |
131 | pjrtErrorCodeFromCInt : Int -> PjrtErrorCode
132 | pjrtErrorCodeFromCInt = \case
133 |   1  => PJRT_Error_Code_CANCELLED
134 |   2  => PJRT_Error_Code_UNKNOWN
135 |   3  => PJRT_Error_Code_INVALID_ARGUMENT
136 |   4  => PJRT_Error_Code_DEADLINE_EXCEEDED
137 |   5  => PJRT_Error_Code_NOT_FOUND
138 |   6  => PJRT_Error_Code_ALREADY_EXISTS
139 |   7  => PJRT_Error_Code_PERMISSION_DENIED
140 |   8  => PJRT_Error_Code_RESOURCE_EXHAUSTED
141 |   9  => PJRT_Error_Code_FAILED_PRECONDITION
142 |   10 => PJRT_Error_Code_ABORTED
143 |   11 => PJRT_Error_Code_OUT_OF_RANGE
144 |   12 => PJRT_Error_Code_UNIMPLEMENTED
145 |   13 => PJRT_Error_Code_INTERNAL
146 |   14 => PJRT_Error_Code_UNAVAILABLE
147 |   15 => PJRT_Error_Code_DATA_LOSS
148 |   16 => PJRT_Error_Code_UNAUTHENTICATED
149 |   n  => assert_total $ idris_crash
150 |     "Unexpected PJRT_Error_Code value received through FFI from XLA: \{show n}"
151 |
152 | ||| A `Pjrt a` produces an `a` or an error from the PJRT layer.
153 | public export 0
154 | Pjrt : Type -> Type
155 | Pjrt = EitherT PjrtError IO
156 |
157 | try : AnyPtr -> AnyPtr -> a -> Pjrt a
158 | try api err onOk = if (isNullPtr err) then right onOk else do
159 |   msg <- pjrtErrorMessage api err
160 |   args <- primIO $ prim__mkPjrtErrorGetCodeArgs err
161 |   getCodeErr <- primIO $ prim__pjrtErrorGetCode api args
162 |   code <- if (isNullPtr getCodeErr) then pure Nothing else do
163 |     let code = prim__pjrtErrorGetCodeArgsCode args
164 |     destroyPjrtError api getCodeErr
165 |     pure $ Just code
166 |   primIO $ prim__deletePjrtErrorGetCodeArgs args
167 |   destroyPjrtError api err
168 |   left $ MkPjrtError msg $ map pjrtErrorCodeFromCInt code
169 |
170 | ||| For internal spidr use only.
171 | export
172 | data PjrtEvent = MkPjrtEvent AnyPtr
173 |
174 | %foreign (ffi "PJRT_Event_Destroy_Args_delete")
175 | prim__deletePjrtEventDestroyArgs : AnyPtr -> PrimIO ()
176 |
177 | %foreign (ffi "PJRT_Event_Destroy_Args_new")
178 | prim__mkPjrtEventDestroyArgs : AnyPtr -> PrimIO AnyPtr
179 |
180 | %foreign (ffi "pjrt_event_destroy")
181 | prim__pjrtEventDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
182 |
183 | %foreign (ffi "PJRT_Event_Await_Args_delete")
184 | prim__deletePjrtEventAwaitArgs : AnyPtr -> PrimIO ()
185 |
186 | %foreign (ffi "PJRT_Event_Await_Args_new")
187 | prim__mkPjrtEventAwaitArgs : AnyPtr -> PrimIO AnyPtr
188 |
189 | %foreign (ffi "pjrt_event_await")
190 | prim__pjrtEventAwait : AnyPtr -> AnyPtr -> PrimIO AnyPtr
191 |
192 | ||| For internal spidr use only.
193 | export
194 | pjrtEventAwait : PjrtApi -> PjrtEvent -> Pjrt ()
195 | pjrtEventAwait (MkPjrtApi api) (MkPjrtEvent event) = do
196 |   args <- primIO $ prim__mkPjrtEventAwaitArgs event
197 |   err <- primIO $ prim__pjrtEventAwait api args
198 |   primIO $ prim__deletePjrtEventAwaitArgs args
199 |   try api err ()
200 |
201 | ||| For use by plugin developers.
202 | export
203 | data PjrtClient = MkPjrtClient GCAnyPtr
204 |
205 | %foreign (ffi "PJRT_Client_Create_Args_delete")
206 | prim__deletePjrtClientCreateArgs : AnyPtr -> PrimIO ()
207 |
208 | %foreign (ffi "PJRT_Client_Create_Args_new")
209 | prim__mkPjrtClientCreateArgs : PrimIO AnyPtr
210 |
211 | %foreign (ffi "PJRT_Client_Create_Args_client")
212 | prim__pjrtClientCreateArgsClient : AnyPtr -> AnyPtr
213 |
214 | %foreign (ffi "pjrt_client_create")
215 | prim__pjrtClientCreate : AnyPtr -> AnyPtr -> PrimIO AnyPtr
216 |
217 | %foreign (ffi "PJRT_Client_Destroy_Args_delete")
218 | prim__deletePjrtClientDestroyArgs : AnyPtr -> PrimIO ()
219 |
220 | %foreign (ffi "PJRT_Client_Destroy_Args_new")
221 | prim__mkPjrtClientDestroyArgs : AnyPtr -> PrimIO AnyPtr
222 |
223 | %foreign (ffi "pjrt_client_destroy")
224 | prim__pjrtClientDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
225 |
226 | handleErrOnDestroy : HasIO io => AnyPtr -> AnyPtr -> String -> io ()
227 | handleErrOnDestroy api err target = unless (isNullPtr err) $ do
228 |   msg <- pjrtErrorMessage api err
229 |   args <- primIO $ prim__mkPjrtErrorGetCodeArgs err
230 |   getCodeErr <- primIO $ prim__pjrtErrorGetCode api args
231 |   if (isNullPtr getCodeErr) then do
232 |       let code = pjrtErrorCodeFromCInt $ prim__pjrtErrorGetCodeArgsCode args
233 |       printLn "WARN: Failed to destroy \{target} with error code \{show code}; message: \{msg}"
234 |     else do
235 |       printLn "WARN: Failed to fetch error code"
236 |       printLn "WARN: Failed to destroy \{target} with unknown error code; message: \{msg}"
237 |       destroyPjrtError api getCodeErr
238 |   primIO $ prim__deletePjrtErrorGetCodeArgs args
239 |   destroyPjrtError api err
240 |
241 | ||| For use by plugin developers.
242 | |||
243 | ||| Create a `PjrtClient`.
244 | export
245 | pjrtClientCreate : PjrtApi -> Pjrt PjrtClient
246 | pjrtClientCreate (MkPjrtApi api) = do
247 |   args <- primIO prim__mkPjrtClientCreateArgs
248 |   err <- primIO $ prim__pjrtClientCreate api args
249 |   let client = prim__pjrtClientCreateArgsClient args
250 |   primIO $ prim__deletePjrtClientCreateArgs args
251 |   try api err =<< do
252 |     client <- onCollectAny' client destroy
253 |     pure $ MkPjrtClient client
254 |
255 |     where
256 |
257 |     destroy : AnyPtr -> IO ()
258 |     destroy client = do
259 |       args <- primIO $ prim__mkPjrtClientDestroyArgs client
260 |       err <- primIO $ prim__pjrtClientDestroy api args
261 |       primIO $ prim__deletePjrtClientDestroyArgs args
262 |       handleErrOnDestroy api err "PJRT_Client"
263 |
264 | ||| For internal spidr use only.
265 | export
266 | data PjrtProgram = MkPjrtProgram AnyPtr
267 |
268 | %foreign (ffi "PJRT_Program_delete")
269 | prim__deletePjrtProgram : AnyPtr -> PrimIO ()
270 |
271 | namespace PjrtProgram
272 |   export
273 |   delete : HasIO io => PjrtProgram -> io ()
274 |   delete (MkPjrtProgram p) = primIO $ prim__deletePjrtProgram p
275 |
276 | %foreign (ffi "PJRT_Program_new")
277 | prim__mkPjrtProgram : Ptr Char -> Bits64 -> PrimIO AnyPtr
278 |
279 | ||| For internal spidr use only.
280 | |||
281 | ||| The `CppString` must live as long as the `PjrtProgram`.
282 | ||| It is up to the caller to deallocate the `PjrtProgram`.
283 | export
284 | mkPjrtProgram : HasIO io => CppString -> io PjrtProgram
285 | mkPjrtProgram (MkCppString code) = MkPjrtProgram <$> (
286 |     primIO $ prim__mkPjrtProgram (prim__stringData code) (prim__stringSize code)
287 |   )
288 |
289 | %foreign (ffi "PJRT_Client_Compile_Args_delete")
290 | prim__deletePjrtClientCompileArgs : AnyPtr -> PrimIO ()
291 |
292 | %foreign (ffi "PJRT_Client_Compile_Args_new")
293 | prim__mkPjrtClientCompileArgs : GCAnyPtr -> AnyPtr -> Ptr Char -> Bits64 -> PrimIO AnyPtr
294 |
295 | %foreign (ffi "PJRT_Client_Compile_Args_executable")
296 | prim__pjrtClientCompileArgsExecutable : AnyPtr -> AnyPtr
297 |
298 | %foreign (ffi "pjrt_client_compile")
299 | prim__pjrtClientCompile : AnyPtr -> AnyPtr -> PrimIO AnyPtr
300 |
301 | %foreign (ffi "PJRT_LoadedExecutable_Destroy_Args_delete")
302 | prim__deletePjrtLoadedExecutableDestroyArgs : AnyPtr -> PrimIO ()
303 |
304 | %foreign (ffi "PJRT_LoadedExecutable_Destroy_Args_new")
305 | prim__mkPjrtLoadedExecutableDestroyArgs : AnyPtr -> PrimIO AnyPtr
306 |
307 | %foreign (ffi "pjrt_loadedexecutable_destroy")
308 | prim__pjrtLoadedExecutableDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
309 |
310 | ||| For internal spidr use only.
311 | export
312 | data PjrtLoadedExecutable = MkPjrtLoadedExecutable AnyPtr
313 |
314 | ||| For internal spidr use only.
315 | export
316 | pjrtLoadedExecutableDestroy : HasIO io => PjrtApi -> PjrtLoadedExecutable -> io ()
317 | pjrtLoadedExecutableDestroy (MkPjrtApi api) (MkPjrtLoadedExecutable executable) = do
318 |   args <- primIO $ prim__mkPjrtLoadedExecutableDestroyArgs executable
319 |   err <- primIO $ prim__pjrtLoadedExecutableDestroy api args
320 |   primIO $ prim__deletePjrtLoadedExecutableDestroyArgs args
321 |   handleErrOnDestroy api err "PJRT_LoadedExecutable"
322 |
323 | ||| For internal spidr use only.
324 | |||
325 | ||| It is up to the caller to deallocate the `PjrtLoadedExecutable`.
326 | export
327 | pjrtClientCompile : PjrtApi -> PjrtClient -> PjrtProgram -> CppString -> Pjrt PjrtLoadedExecutable
328 | pjrtClientCompile
329 |   (MkPjrtApi api) (MkPjrtClient client) (MkPjrtProgram program) (MkCppString options) = do
330 |     args <- primIO $ prim__mkPjrtClientCompileArgs
331 |       client program (prim__stringData options) (prim__stringSize options)
332 |     err <- primIO $ prim__pjrtClientCompile api args
333 |     let executable = prim__pjrtClientCompileArgsExecutable args
334 |     primIO $ prim__deletePjrtClientCompileArgs args
335 |     try api err $ MkPjrtLoadedExecutable executable
336 |
337 | %foreign (ffi "PJRT_ExecuteOptions_delete")
338 | prim__deletePjrtExecuteOptions : AnyPtr -> PrimIO ()
339 |
340 | %foreign (ffi "PJRT_ExecuteOptions_new")
341 | prim__mkPjrtExecuteOptions : PrimIO AnyPtr
342 |
343 | %foreign (ffi "PJRT_LoadedExecutable_Execute_Args_delete")
344 | prim__deletePjrtLoadedExecutableExecuteArgs : AnyPtr -> PrimIO ()
345 |
346 | %foreign (ffi "PJRT_LoadedExecutable_Execute_Args_new")
347 | prim__mkPjrtLoadedExecutableExecuteArgs : AnyPtr -> AnyPtr -> AnyPtr -> PrimIO AnyPtr
348 |
349 | %foreign (ffi "pjrt_loadedexecutable_execute")
350 | prim__pjrtLoadedExecutableExecute : AnyPtr -> AnyPtr -> PrimIO AnyPtr
351 |
352 | %foreign (ffi "PJRT_Buffer_Destroy_Args_delete")
353 | prim__deletePjrtBufferDestroyArgs : AnyPtr -> PrimIO ()
354 |
355 | %foreign (ffi "PJRT_Buffer_Destroy_Args_new")
356 | prim__mkPjrtBufferDestroyArgs : AnyPtr -> PrimIO AnyPtr
357 |
358 | %foreign (ffi "pjrt_buffer_destroy")
359 | prim__pjrtBufferDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
360 |
361 | ||| For internal spidr use only.
362 | export
363 | data PjrtBuffer = MkPjrtBuffer AnyPtr
364 |
365 | ||| For internal spidr use only.
366 | export
367 | pjrtBufferDestroy : HasIO io => PjrtApi -> PjrtBuffer -> io ()
368 | pjrtBufferDestroy (MkPjrtApi api) (MkPjrtBuffer buffer) = do
369 |   args <- primIO $ prim__mkPjrtBufferDestroyArgs buffer
370 |   err <- primIO $ prim__pjrtBufferDestroy api args
371 |   primIO $ prim__deletePjrtBufferDestroyArgs args
372 |   handleErrOnDestroy api err "PJRT_Buffer"
373 |
374 | ||| For internal spidr use only.
375 | |||
376 | ||| It is up to the caller to deallocate the `PjrtBuffer`s.
377 | export
378 | pjrtLoadedExecutableExecute :
379 |   PjrtApi -> PjrtLoadedExecutable -> (outputs : Nat) -> Pjrt (Vect outputs PjrtBuffer)
380 | pjrtLoadedExecutableExecute (MkPjrtApi api) (MkPjrtLoadedExecutable executable) outputs = do
381 |   outputListsInner <- malloc (cast outputs * cast sizeofPtr)
382 |   outputLists <- malloc $ cast sizeofPtr
383 |   primIO $ prim__setArrayVoidPtr outputLists 0 outputListsInner
384 |   options <- primIO prim__mkPjrtExecuteOptions
385 |   args <- primIO $ prim__mkPjrtLoadedExecutableExecuteArgs executable options outputLists
386 |   err <- primIO $ prim__pjrtLoadedExecutableExecute api args
387 |   primIO $ prim__deletePjrtLoadedExecutableExecuteArgs args
388 |   primIO $ prim__deletePjrtExecuteOptions options
389 |   let buffers = map (\o => MkPjrtBuffer $ prim__getArrayVoidPtr outputListsInner (cast o)) (range outputs)
390 |   free outputLists
391 |   free outputListsInner
392 |   try api err buffers
393 |
394 | %foreign (ffi "PJRT_Buffer_ToHostBuffer_Args_delete")
395 | prim__deletePjrtBufferToHostBufferArgs : AnyPtr -> PrimIO ()
396 |
397 | %foreign (ffi "PJRT_Buffer_ToHostBuffer_Args_new")
398 | prim__mkPjrtBufferToHostBufferArgs : AnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
399 |
400 | %foreign (ffi "PJRT_Buffer_ToHostBuffer_Args_event")
401 | prim__pjrtBufferToHostBufferArgsEvent : AnyPtr -> AnyPtr
402 |
403 | %foreign (ffi "pjrt_buffer_tohostbuffer")
404 | prim__pjrtBufferToHostBuffer : AnyPtr -> AnyPtr -> PrimIO AnyPtr
405 |
406 | ||| For internal spidr use only.
407 | export
408 | pjrtEventDestroy : HasIO io => PjrtApi -> PjrtEvent -> io ()
409 | pjrtEventDestroy (MkPjrtApi api) (MkPjrtEvent event) = do
410 |   args <- primIO $ prim__mkPjrtEventDestroyArgs event
411 |   err <- primIO $ prim__pjrtEventDestroy api args
412 |   primIO $ prim__deletePjrtEventDestroyArgs args
413 |   handleErrOnDestroy api err "PJRT_Event"
414 |
415 | ||| For internal spidr use only.
416 | |||
417 | ||| This is a synchronous variant of PJRT's `PJRT_Buffer_ToHostBuffer`.
418 | ||| Unlike the PJRT API, it also handles zero-byte arrays.
419 | export
420 | pjrtBufferToHostBuffer : PjrtApi -> PjrtBuffer -> ArrayType a => Array a -> Pjrt ()
421 | pjrtBufferToHostBuffer (MkPjrtApi api) (MkPjrtBuffer buffer) (MkArray arr arrSize) =
422 |   when (arrSize > 0) $ do
423 |     args <- primIO $ prim__mkPjrtBufferToHostBufferArgs buffer arr (arrSize * elemSize {a})
424 |     err <- primIO $ prim__pjrtBufferToHostBuffer api args
425 |     let event = MkPjrtEvent $ prim__pjrtBufferToHostBufferArgsEvent args
426 |     primIO $ prim__deletePjrtBufferToHostBufferArgs args
427 |     try api err ()
428 |     let api = MkPjrtApi api
429 |     pjrtEventAwait api event
430 |     pjrtEventDestroy api event
431 |