19 | module Compiler.Xla.PJRT.C.PjrtCApi
21 | import public Control.Monad.Either
22 | import Derive.Prelude
23 | import Language.Reflection
26 | import Compiler.String
27 | import Compiler.Array
31 | %language ElabReflection
33 | ffi : String -> String
34 | ffi = libxla "c/xla/pjrt/c/pjrt_c_api.h"
41 | data PjrtApi = MkPjrtApi AnyPtr
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
63 | %runElab derive "PjrtErrorCode" [Show]
67 | record PjrtError where
68 | constructor MkPjrtError
74 | code : Maybe PjrtErrorCode
77 | Show PjrtError where
79 | let code = case e.code of
80 | Nothing => "unknown"
82 | in "PjrtError (error code \{code})\n\{e.message}"
84 | %foreign (ffi "PJRT_Error_Destroy_Args_delete")
85 | prim__deletePjrtErrorDestroyArgs : AnyPtr -> PrimIO ()
87 | %foreign (ffi "PJRT_Error_Destroy_Args_new")
88 | prim__mkPjrtErrorDestroyArgs : AnyPtr -> PrimIO AnyPtr
90 | %foreign (ffi "pjrt_error_destroy")
91 | prim__pjrtErrorDestroy : AnyPtr -> AnyPtr -> PrimIO ()
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
99 | %foreign (ffi "PJRT_Error_Message_Args_delete")
100 | prim__deletePjrtErrorMessageArgs : AnyPtr -> PrimIO ()
102 | %foreign (ffi "PJRT_Error_Message_Args_new")
103 | prim__mkPjrtErrorMessageArgs : AnyPtr -> PrimIO AnyPtr
105 | %foreign (ffi "PJRT_Error_Message_Args_message")
106 | prim__pjrtErrorMessageArgsMessage : AnyPtr -> PrimIO String
108 | %foreign (ffi "pjrt_error_message")
109 | prim__pjrtErrorMessage : AnyPtr -> AnyPtr -> PrimIO ()
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
119 | %foreign (ffi "PJRT_Error_GetCode_Args_delete")
120 | prim__deletePjrtErrorGetCodeArgs : AnyPtr -> PrimIO ()
122 | %foreign (ffi "PJRT_Error_GetCode_Args_new")
123 | prim__mkPjrtErrorGetCodeArgs : AnyPtr -> PrimIO AnyPtr
125 | %foreign (ffi "PJRT_Error_GetCode_Args_code")
126 | prim__pjrtErrorGetCodeArgsCode : AnyPtr -> Int
128 | %foreign (ffi "pjrt_error_getcode")
129 | prim__pjrtErrorGetCode : AnyPtr -> AnyPtr -> PrimIO AnyPtr
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}"
154 | Pjrt : Type -> Type
155 | Pjrt = EitherT PjrtError IO
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
166 | primIO $
prim__deletePjrtErrorGetCodeArgs args
167 | destroyPjrtError api err
168 | left $
MkPjrtError msg $
map pjrtErrorCodeFromCInt code
172 | data PjrtEvent = MkPjrtEvent AnyPtr
174 | %foreign (ffi "PJRT_Event_Destroy_Args_delete")
175 | prim__deletePjrtEventDestroyArgs : AnyPtr -> PrimIO ()
177 | %foreign (ffi "PJRT_Event_Destroy_Args_new")
178 | prim__mkPjrtEventDestroyArgs : AnyPtr -> PrimIO AnyPtr
180 | %foreign (ffi "pjrt_event_destroy")
181 | prim__pjrtEventDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
183 | %foreign (ffi "PJRT_Event_Await_Args_delete")
184 | prim__deletePjrtEventAwaitArgs : AnyPtr -> PrimIO ()
186 | %foreign (ffi "PJRT_Event_Await_Args_new")
187 | prim__mkPjrtEventAwaitArgs : AnyPtr -> PrimIO AnyPtr
189 | %foreign (ffi "pjrt_event_await")
190 | prim__pjrtEventAwait : AnyPtr -> AnyPtr -> PrimIO AnyPtr
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
203 | data PjrtClient = MkPjrtClient GCAnyPtr
205 | %foreign (ffi "PJRT_Client_Create_Args_delete")
206 | prim__deletePjrtClientCreateArgs : AnyPtr -> PrimIO ()
208 | %foreign (ffi "PJRT_Client_Create_Args_new")
209 | prim__mkPjrtClientCreateArgs : PrimIO AnyPtr
211 | %foreign (ffi "PJRT_Client_Create_Args_client")
212 | prim__pjrtClientCreateArgsClient : AnyPtr -> AnyPtr
214 | %foreign (ffi "pjrt_client_create")
215 | prim__pjrtClientCreate : AnyPtr -> AnyPtr -> PrimIO AnyPtr
217 | %foreign (ffi "PJRT_Client_Destroy_Args_delete")
218 | prim__deletePjrtClientDestroyArgs : AnyPtr -> PrimIO ()
220 | %foreign (ffi "PJRT_Client_Destroy_Args_new")
221 | prim__mkPjrtClientDestroyArgs : AnyPtr -> PrimIO AnyPtr
223 | %foreign (ffi "pjrt_client_destroy")
224 | prim__pjrtClientDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
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}"
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
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
252 | client <- onCollectAny' client destroy
253 | pure $
MkPjrtClient client
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"
266 | data PjrtProgram = MkPjrtProgram AnyPtr
268 | %foreign (ffi "PJRT_Program_delete")
269 | prim__deletePjrtProgram : AnyPtr -> PrimIO ()
271 | namespace PjrtProgram
273 | delete : HasIO io => PjrtProgram -> io ()
274 | delete (MkPjrtProgram p) = primIO $
prim__deletePjrtProgram p
276 | %foreign (ffi "PJRT_Program_new")
277 | prim__mkPjrtProgram : Ptr Char -> Bits64 -> PrimIO AnyPtr
284 | mkPjrtProgram : HasIO io => CppString -> io PjrtProgram
285 | mkPjrtProgram (MkCppString code) = MkPjrtProgram <$> (
286 | primIO $
prim__mkPjrtProgram (prim__stringData code) (prim__stringSize code)
289 | %foreign (ffi "PJRT_Client_Compile_Args_delete")
290 | prim__deletePjrtClientCompileArgs : AnyPtr -> PrimIO ()
292 | %foreign (ffi "PJRT_Client_Compile_Args_new")
293 | prim__mkPjrtClientCompileArgs : GCAnyPtr -> AnyPtr -> Ptr Char -> Bits64 -> PrimIO AnyPtr
295 | %foreign (ffi "PJRT_Client_Compile_Args_executable")
296 | prim__pjrtClientCompileArgsExecutable : AnyPtr -> AnyPtr
298 | %foreign (ffi "pjrt_client_compile")
299 | prim__pjrtClientCompile : AnyPtr -> AnyPtr -> PrimIO AnyPtr
301 | %foreign (ffi "PJRT_LoadedExecutable_Destroy_Args_delete")
302 | prim__deletePjrtLoadedExecutableDestroyArgs : AnyPtr -> PrimIO ()
304 | %foreign (ffi "PJRT_LoadedExecutable_Destroy_Args_new")
305 | prim__mkPjrtLoadedExecutableDestroyArgs : AnyPtr -> PrimIO AnyPtr
307 | %foreign (ffi "pjrt_loadedexecutable_destroy")
308 | prim__pjrtLoadedExecutableDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
312 | data PjrtLoadedExecutable = MkPjrtLoadedExecutable AnyPtr
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"
327 | pjrtClientCompile : PjrtApi -> PjrtClient -> PjrtProgram -> CppString -> Pjrt PjrtLoadedExecutable
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
337 | %foreign (ffi "PJRT_ExecuteOptions_delete")
338 | prim__deletePjrtExecuteOptions : AnyPtr -> PrimIO ()
340 | %foreign (ffi "PJRT_ExecuteOptions_new")
341 | prim__mkPjrtExecuteOptions : PrimIO AnyPtr
343 | %foreign (ffi "PJRT_LoadedExecutable_Execute_Args_delete")
344 | prim__deletePjrtLoadedExecutableExecuteArgs : AnyPtr -> PrimIO ()
346 | %foreign (ffi "PJRT_LoadedExecutable_Execute_Args_new")
347 | prim__mkPjrtLoadedExecutableExecuteArgs : AnyPtr -> AnyPtr -> AnyPtr -> PrimIO AnyPtr
349 | %foreign (ffi "pjrt_loadedexecutable_execute")
350 | prim__pjrtLoadedExecutableExecute : AnyPtr -> AnyPtr -> PrimIO AnyPtr
352 | %foreign (ffi "PJRT_Buffer_Destroy_Args_delete")
353 | prim__deletePjrtBufferDestroyArgs : AnyPtr -> PrimIO ()
355 | %foreign (ffi "PJRT_Buffer_Destroy_Args_new")
356 | prim__mkPjrtBufferDestroyArgs : AnyPtr -> PrimIO AnyPtr
358 | %foreign (ffi "pjrt_buffer_destroy")
359 | prim__pjrtBufferDestroy : AnyPtr -> AnyPtr -> PrimIO AnyPtr
363 | data PjrtBuffer = MkPjrtBuffer AnyPtr
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"
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)
391 | free outputListsInner
392 | try api err buffers
394 | %foreign (ffi "PJRT_Buffer_ToHostBuffer_Args_delete")
395 | prim__deletePjrtBufferToHostBufferArgs : AnyPtr -> PrimIO ()
397 | %foreign (ffi "PJRT_Buffer_ToHostBuffer_Args_new")
398 | prim__mkPjrtBufferToHostBufferArgs : AnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
400 | %foreign (ffi "PJRT_Buffer_ToHostBuffer_Args_event")
401 | prim__pjrtBufferToHostBufferArgsEvent : AnyPtr -> AnyPtr
403 | %foreign (ffi "pjrt_buffer_tohostbuffer")
404 | prim__pjrtBufferToHostBuffer : AnyPtr -> AnyPtr -> PrimIO AnyPtr
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"
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
428 | let api = MkPjrtApi api
429 | pjrtEventAwait api event
430 | pjrtEventDestroy api event