jon.recoil.org

Source file onnxrt.ml

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
module Dtype = struct
  type ('ocaml, 'elt) t =
    | Float32 : (float, Bigarray.float32_elt) t
    | Float64 : (float, Bigarray.float64_elt) t
    | Int8 : (int, Bigarray.int8_signed_elt) t
    | Uint8 : (int, Bigarray.int8_unsigned_elt) t
    | Int16 : (int, Bigarray.int16_signed_elt) t
    | Uint16 : (int, Bigarray.int16_unsigned_elt) t
    | Int32 : (int32, Bigarray.int32_elt) t

  type packed = Pack : ('ocaml, 'elt) t -> packed

  let to_string : type a b. (a, b) t -> string = function
    | Float32 -> "float32"
    | Float64 -> "float64"
    | Int8 -> "int8"
    | Uint8 -> "uint8"
    | Int16 -> "int16"
    | Uint16 -> "uint16"
    | Int32 -> "int32"

  let of_string = function
    | "float32" -> Some (Pack Float32)
    | "float64" -> Some (Pack Float64)
    | "int8" -> Some (Pack Int8)
    | "uint8" -> Some (Pack Uint8)
    | "int16" -> Some (Pack Int16)
    | "uint16" -> Some (Pack Uint16)
    | "int32" -> Some (Pack Int32)
    | _ -> None

  let equal : type a b c d. (a, b) t -> (c, d) t -> bool =
   fun a b ->
    match (a, b) with
    | Float32, Float32 -> true
    | Float64, Float64 -> true
    | Int8, Int8 -> true
    | Uint8, Uint8 -> true
    | Int16, Int16 -> true
    | Uint16, Uint16 -> true
    | Int32, Int32 -> true
    | _ -> false

  let _to_bigarray_kind : type a b. (a, b) t -> (a, b) Bigarray.kind = function
    | Float32 -> Bigarray.float32
    | Float64 -> Bigarray.float64
    | Int8 -> Bigarray.int8_signed
    | Uint8 -> Bigarray.int8_unsigned
    | Int16 -> Bigarray.int16_signed
    | Uint16 -> Bigarray.int16_unsigned
    | Int32 -> Bigarray.int32
end

module Tensor = struct
  type t = {
    js_tensor : Js_of_ocaml.Js.Unsafe.any;
    mutable disposed : bool;
  }

  type location = Cpu | Gpu_buffer

  let check_not_disposed t =
    if t.disposed then invalid_arg "Tensor has been disposed"

  let check_cpu t =
    check_not_disposed t;
    let loc = Js_helpers.get_string
      (Js_of_ocaml.Js.Unsafe.coerce t.js_tensor) "location" in
    if loc <> "cpu" then
      invalid_arg "Tensor data is on GPU; use Tensor.download first"

  let check_dtype : type a b. (a, b) Dtype.t -> t -> unit =
   fun expected t ->
    let actual_str = Js_helpers.get_string
      (Js_of_ocaml.Js.Unsafe.coerce t.js_tensor) "type" in
    let expected_str = Dtype.to_string expected in
    if actual_str <> expected_str then
      failwith (Printf.sprintf "Dtype mismatch: tensor is %s, expected %s"
                  actual_str expected_str)

  let typed_array_of_bigarray :
    type a b. (a, b) Dtype.t ->
    (a, b, Bigarray.c_layout) Bigarray.Array1.t ->
    Js_of_ocaml.Js.Unsafe.any =
   fun dtype ba ->
    let open Js_of_ocaml in
    let ga = Bigarray.genarray_of_array1 ba in
    match dtype with
    | Dtype.Float32 ->
        Js.Unsafe.coerce (Typed_array.from_genarray Typed_array.Float32 ga)
    | Dtype.Float64 ->
        Js.Unsafe.coerce (Typed_array.from_genarray Typed_array.Float64 ga)
    | Dtype.Int8 ->
        Js.Unsafe.coerce (Typed_array.from_genarray Typed_array.Int8_signed ga)
    | Dtype.Uint8 ->
        Js.Unsafe.coerce (Typed_array.from_genarray Typed_array.Int8_unsigned ga)
    | Dtype.Int16 ->
        Js.Unsafe.coerce (Typed_array.from_genarray Typed_array.Int16_signed ga)
    | Dtype.Uint16 ->
        Js.Unsafe.coerce (Typed_array.from_genarray Typed_array.Int16_unsigned ga)
    | Dtype.Int32 ->
        Js.Unsafe.coerce (Typed_array.from_genarray Typed_array.Int32_signed ga)

  let of_bigarray1 :
    type a b. (a, b) Dtype.t ->
    (a, b, Bigarray.c_layout) Bigarray.Array1.t ->
    dims:int array -> t =
   fun dtype ba ~dims ->
    let expected_size = Array.fold_left ( * ) 1 dims in
    let actual_size = Bigarray.Array1.dim ba in
    if expected_size <> actual_size then
      invalid_arg (Printf.sprintf
        "Tensor.of_bigarray1: dims product (%d) <> bigarray length (%d)"
        expected_size actual_size);
    let open Js_of_ocaml in
    let ta = typed_array_of_bigarray dtype ba in
    let js_tensor =
      Js.Unsafe.new_obj
        (Js.Unsafe.get (Js_helpers.ort ()) (Js.string "Tensor"))
        [| Js.Unsafe.inject (Js.string (Dtype.to_string dtype));
           ta;
           Js.Unsafe.inject (Js_helpers.js_int_array dims) |]
    in
    { js_tensor = Js.Unsafe.coerce js_tensor; disposed = false }

  let of_bigarray :
    type a b. (a, b) Dtype.t ->
    (a, b, Bigarray.c_layout) Bigarray.Genarray.t -> t =
   fun dtype ga ->
    let dims = Bigarray.Genarray.dims ga in
    let flat = Bigarray.reshape_1 ga (Array.fold_left ( * ) 1 dims) in
    of_bigarray1 dtype flat ~dims

  let of_float32s data ~dims =
    let expected_size = Array.fold_left ( * ) 1 dims in
    if Array.length data <> expected_size then
      invalid_arg (Printf.sprintf
        "Tensor.of_float32s: array length (%d) <> dims product (%d)"
        (Array.length data) expected_size);
    let ba = Bigarray.Array1.create Bigarray.float32 Bigarray.c_layout
               expected_size in
    Array.iteri (fun i v -> Bigarray.Array1.set ba i v) data;
    of_bigarray1 Float32 ba ~dims

  let to_bigarray1_exn :
    type a b. (a, b) Dtype.t -> t ->
    (a, b, Bigarray.c_layout) Bigarray.Array1.t =
   fun dtype t ->
    check_cpu t;
    check_dtype dtype t;
    let open Js_of_ocaml in
    let data : Js.Unsafe.any = Js.Unsafe.get t.js_tensor (Js.string "data") in
    let ta = (Js.Unsafe.coerce data : (_, _, _) Typed_array.typedArray Js.t) in
    let ga = Typed_array.to_genarray ta in
    let size = Bigarray.Genarray.nth_dim ga 0 in
    let ba = Bigarray.reshape_1 ga size in
    (Obj.magic ba : (a, b, Bigarray.c_layout) Bigarray.Array1.t)

  let to_bigarray_exn :
    type a b. (a, b) Dtype.t -> t ->
    (a, b, Bigarray.c_layout) Bigarray.Genarray.t =
   fun dtype t ->
    let flat = to_bigarray1_exn dtype t in
    let dims_js : Js_of_ocaml.Js.Unsafe.any =
      Js_of_ocaml.Js.Unsafe.get t.js_tensor (Js_of_ocaml.Js.string "dims") in
    let dims = Js_helpers.int_array_of_js (Js_of_ocaml.Js.Unsafe.coerce dims_js) in
    Bigarray.genarray_of_array1 flat |> fun ga -> Bigarray.reshape ga dims

  let download :
    type a b. (a, b) Dtype.t -> t ->
    (a, b, Bigarray.c_layout) Bigarray.Array1.t Lwt.t =
   fun dtype t ->
    check_not_disposed t;
    check_dtype dtype t;
    let open Js_of_ocaml in
    let promise = Js.Unsafe.meth_call t.js_tensor "getData" [||] in
    let open Lwt.Syntax in
    let+ data = Promise_lwt.to_lwt promise in
    let ta = (Js.Unsafe.coerce data : (_, _, _) Typed_array.typedArray Js.t) in
    let ga = Typed_array.to_genarray ta in
    let size = Bigarray.Genarray.nth_dim ga 0 in
    let ba = Bigarray.reshape_1 ga size in
    (Obj.magic ba : (a, b, Bigarray.c_layout) Bigarray.Array1.t)

  let dims t =
    check_not_disposed t;
    let open Js_of_ocaml in
    let dims_js = Js.Unsafe.get t.js_tensor (Js.string "dims") in
    Js_helpers.int_array_of_js (Js.Unsafe.coerce dims_js)

  let dtype t =
    check_not_disposed t;
    let type_str = Js_helpers.get_string
      (Js_of_ocaml.Js.Unsafe.coerce t.js_tensor) "type" in
    match Dtype.of_string type_str with
    | Some p -> p
    | None -> failwith (Printf.sprintf "Unknown tensor dtype: %s" type_str)

  let size t =
    check_not_disposed t;
    Js_helpers.get_int (Js_of_ocaml.Js.Unsafe.coerce t.js_tensor) "size"

  let location t =
    check_not_disposed t;
    let loc = Js_helpers.get_string
      (Js_of_ocaml.Js.Unsafe.coerce t.js_tensor) "location" in
    match loc with
    | "cpu" -> Cpu
    | "gpu-buffer" -> Gpu_buffer
    | s -> failwith (Printf.sprintf "Unknown tensor location: %s" s)

  let dispose t =
    if not t.disposed then begin
      let open Js_of_ocaml in
      ignore (Js.Unsafe.meth_call t.js_tensor "dispose" [||] : Js.Unsafe.any);
      t.disposed <- true
    end
end

module Execution_provider = struct
  type t = Wasm | Webgpu
  let to_string = function Wasm -> "wasm" | Webgpu -> "webgpu"
end

type output_location = Cpu | Gpu_buffer
type graph_optimization = Disabled | Basic | Extended | All

let graph_optimization_to_string = function
  | Disabled -> "disabled"
  | Basic -> "basic"
  | Extended -> "extended"
  | All -> "all"

let output_location_to_js = function
  | Cpu -> Js_of_ocaml.Js.string "cpu"
  | Gpu_buffer -> Js_of_ocaml.Js.string "gpu-buffer"

let log_level_to_string = function
  | `Verbose -> "verbose"
  | `Info -> "info"
  | `Warning -> "warning"
  | `Error -> "error"
  | `Fatal -> "fatal"

module Session = struct
  type t = {
    js_session : Js_of_ocaml.Js.Unsafe.any;
    input_names_ : string list;
    output_names_ : string list;
  }

  let build_options ?execution_providers ?graph_optimization
      ?preferred_output_location ?log_level () =
    let open Js_of_ocaml in
    let pairs = ref [] in
    (match execution_providers with
     | Some eps ->
         let js_eps = Js.array (Array.of_list
           (List.map (fun ep ->
              Js.Unsafe.inject (Js.string (Execution_provider.to_string ep)))
             eps)) in
         pairs := ("executionProviders", Js.Unsafe.inject js_eps) :: !pairs
     | None -> ());
    (match graph_optimization with
     | Some go ->
         pairs := ("graphOptimizationLevel",
                    Js.Unsafe.inject (Js.string (graph_optimization_to_string go)))
                  :: !pairs
     | None -> ());
    (match preferred_output_location with
     | Some loc ->
         pairs := ("preferredOutputLocation",
                    Js.Unsafe.inject (output_location_to_js loc))
                  :: !pairs
     | None -> ());
    (match log_level with
     | Some level ->
         pairs := ("logSeverityLevel",
                    Js.Unsafe.inject (Js.string (log_level_to_string level)))
                  :: !pairs
     | None -> ());
    Js.Unsafe.obj (Array.of_list !pairs)

  let wrap_session js_session =
    let open Js_of_ocaml in
    let input_names_ = Js_helpers.string_list_of_js_array
      (Js.Unsafe.coerce (Js.Unsafe.get js_session (Js.string "inputNames"))) in
    let output_names_ = Js_helpers.string_list_of_js_array
      (Js.Unsafe.coerce (Js.Unsafe.get js_session (Js.string "outputNames"))) in
    { js_session = Js.Unsafe.coerce js_session; input_names_; output_names_ }

  let create ?execution_providers ?graph_optimization
      ?preferred_output_location ?log_level model_url () =
    let open Js_of_ocaml in
    let ort = Js_helpers.ort () in
    let inference_session = Js.Unsafe.get ort (Js.string "InferenceSession") in
    let options = build_options ?execution_providers ?graph_optimization
        ?preferred_output_location ?log_level () in
    let promise = Js.Unsafe.meth_call inference_session "create"
        [| Js.Unsafe.inject (Js.string model_url);
           Js.Unsafe.inject options |] in
    let open Lwt.Syntax in
    let+ js_session = Promise_lwt.to_lwt promise in
    wrap_session js_session

  let create_from_buffer (type a b) ?execution_providers ?graph_optimization
      ?preferred_output_location ?log_level
      (buffer : (a, b, Bigarray.c_layout) Bigarray.Array1.t) () =
    let open Js_of_ocaml in
    let ort = Js_helpers.ort () in
    let inference_session = Js.Unsafe.get ort (Js.string "InferenceSession") in
    let options = build_options ?execution_providers ?graph_optimization
        ?preferred_output_location ?log_level () in
    let ga = Bigarray.genarray_of_array1 buffer in
    let ta = Typed_array.from_genarray Typed_array.Int8_unsigned (Obj.magic ga) in
    let ab : Typed_array.arrayBuffer Js.t =
      Js.Unsafe.get (Js.Unsafe.coerce ta) (Js.string "buffer") in
    let uint8 = Js.Unsafe.new_obj
      (Js.Unsafe.global##._Uint8Array)
      [| Js.Unsafe.inject ab |] in
    let promise = Js.Unsafe.meth_call inference_session "create"
        [| Js.Unsafe.inject uint8;
           Js.Unsafe.inject options |] in
    let open Lwt.Syntax in
    let+ js_session = Promise_lwt.to_lwt promise in
    wrap_session js_session

  let run t inputs =
    let open Js_of_ocaml in
    let feeds = Js.Unsafe.obj
      (Array.of_list
         (List.map (fun (name, (tensor : Tensor.t)) ->
            (name, Js.Unsafe.inject tensor.js_tensor))
           inputs)) in
    let promise = Js.Unsafe.meth_call t.js_session "run"
        [| Js.Unsafe.inject feeds |] in
    let open Lwt.Syntax in
    let+ results = Promise_lwt.to_lwt promise in
    List.map (fun name ->
      let js_tensor = Js.Unsafe.get results (Js.string name) in
      (name, Tensor.{ js_tensor = Js.Unsafe.coerce js_tensor;
                       disposed = false }))
      t.output_names_

  let run_with_outputs t inputs ~output_names =
    let open Js_of_ocaml in
    let feeds = Js.Unsafe.obj
      (Array.of_list
         (List.map (fun (name, (tensor : Tensor.t)) ->
            (name, Js.Unsafe.inject tensor.js_tensor))
           inputs)) in
    let promise = Js.Unsafe.meth_call t.js_session "run"
        [| Js.Unsafe.inject feeds |] in
    let open Lwt.Syntax in
    let+ results = Promise_lwt.to_lwt promise in
    List.map (fun name ->
      let js_tensor = Js.Unsafe.get results (Js.string name) in
      (name, Tensor.{ js_tensor = Js.Unsafe.coerce js_tensor;
                       disposed = false }))
      output_names

  let input_names t = t.input_names_
  let output_names t = t.output_names_

  let release t =
    let open Js_of_ocaml in
    let promise = Js.Unsafe.meth_call t.js_session "release" [||] in
    Promise_lwt.to_lwt promise |> Lwt.map (fun (_ : Js.Unsafe.any) -> ())
end

module Env = struct
  module Wasm = struct
    let set_num_threads n =
      let ort = Js_helpers.ort () in
      Js_helpers.set
        (Js_helpers.get_nested ort "env" "wasm")
        "numThreads" n

    let set_simd enabled =
      let ort = Js_helpers.ort () in
      Js_helpers.set
        (Js_helpers.get_nested ort "env" "wasm")
        "simd" (Js_of_ocaml.Js.bool enabled)

    let set_proxy enabled =
      let ort = Js_helpers.ort () in
      Js_helpers.set
        (Js_helpers.get_nested ort "env" "wasm")
        "proxy" (Js_of_ocaml.Js.bool enabled)

    let set_wasm_paths prefix =
      let ort = Js_helpers.ort () in
      Js_helpers.set
        (Js_helpers.get_nested ort "env" "wasm")
        "wasmPaths" (Js_of_ocaml.Js.string prefix)
  end

  module Webgpu = struct
    let set_power_preference pref =
      let ort = Js_helpers.ort () in
      let s = match pref with
        | `High_performance -> "high-performance"
        | `Low_power -> "low-power"
      in
      Js_helpers.set
        (Js_helpers.get_nested ort "env" "webgpu")
        "powerPreference" (Js_of_ocaml.Js.string s)
  end
end