lib/accy/src/tensor/wire/decode.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const accy = @import("../../root.zig");
  3 const tensor = @import("../root.zig");
  4 const format = @import("format.zig");
  5 
  6 const program_mod = tensor.program;
  7 const Builder = tensor.trace.Builder;
  8 const ByteReader = accy.artifact.wire.ByteReader;
  9 const Dim = tensor.Dim;
 10 const Id = program_mod.Id;
 11 const Kind = program_mod.Kind;
 12 const Operation = program_mod.Operation;
 13 const Program = program_mod.Program;
 14 const Tag = format.Tag;
 15 const Type = program_mod.Type;
 16 const Value = tensor.trace.Value;
 17 
 18 const max_operands = program_mod.max_operation_operands;
 19 const min_dim_bytes = @sizeOf(u32) + @sizeOf(i64);
 20 const derived_result = Type.scalar(.i1);
 21 
 22 /// A caller uses this function to rebuild a tensor program from wire bytes, the versioned byte
 23 /// encoding of a tensor program, sent by another process. The function reads version 1 bytes and
 24 /// replays each operation through `Builder.operation` on the program builder, the object that
 25 /// builds a program one checked operation at a time, which works out every result type again and
 26 /// checks it, so the bytes never supply a result type the builder did not derive. Bytes that end
 27 /// early or carry anything after the program return `error.InvalidArtifact`. A wrong magic number
 28 /// returns `error.BadMagic` and another version returns `error.UnsupportedVersion`. The caller owns
 29 /// the returned program and frees it with `Program.deinit`. The builder copies names and payloads
 30 /// into the program's own memory, so the caller may free `bytes` as soon as the call returns.
 31 pub fn decode(allocator: std.mem.Allocator, bytes: []const u8) !Program {
 32     var decoder = Decoder{ .allocator = allocator, .reader = .{ .bytes = bytes } };
 33     defer decoder.deinit();
 34     if (try decoder.reader.readU32() != format.magic) return error.BadMagic;
 35     if (try decoder.reader.readU32() != format.version) return error.UnsupportedVersion;
 36     const name = try decoder.reader.readLengthPrefixedBytes();
 37     try decoder.open(name, .{});
 38     for (0..bytes.len) |_| {
 39         const frame = &decoder.frames[decoder.count - 1];
 40         if (frame.remaining != 0) {
 41             frame.remaining -= 1;
 42             try decoder.operation(frame);
 43         } else if (decoder.count == 1) {
 44             return decoder.finish(frame);
 45         } else {
 46             try decoder.close();
 47         }
 48     }
 49     unreachable;
 50 }
 51 
 52 const Scan = struct {
 53     length: i64 = 0,
 54     inits: [program_mod.max_scan_carries]Id = undefined,
 55     init_count: usize = 0,
 56 };
 57 
 58 const Frame = struct {
 59     builder: Builder,
 60     values: std.ArrayListUnmanaged(Value) = .empty,
 61     remaining: u32,
 62     scan: Scan,
 63 };
 64 
 65 const Span = struct {
 66     start: usize,
 67     end: usize,
 68 };
 69 
 70 const Decoder = struct {
 71     allocator: std.mem.Allocator,
 72     reader: ByteReader,
 73     frames: [format.max_scan_depth + 1]Frame = undefined,
 74     count: usize = 0,
 75     dims: std.ArrayListUnmanaged(Dim) = .empty,
 76     integers: std.ArrayListUnmanaged(i64) = .empty,
 77     operands: [max_operands]Id = undefined,
 78 
 79     fn deinit(self: *Decoder) void {
 80         std.debug.assert(self.count <= self.frames.len);
 81         for (self.frames[0..self.count]) |*frame| {
 82             frame.values.deinit(self.allocator);
 83             frame.builder.deinit();
 84         }
 85         self.integers.deinit(self.allocator);
 86         self.dims.deinit(self.allocator);
 87         self.* = undefined;
 88     }
 89 
 90     fn open(self: *Decoder, name: []const u8, scan: Scan) !void {
 91         std.debug.assert(self.count < self.frames.len);
 92         std.debug.assert(scan.init_count <= scan.inits.len);
 93         const remaining = try self.reader.readU32();
 94         self.frames[self.count] = .{
 95             .builder = try Builder.init(self.allocator, name),
 96             .remaining = remaining,
 97             .scan = scan,
 98         };
 99         self.count += 1;
100     }
101 
102     fn operation(self: *Decoder, frame: *Frame) !void {
103         const tag = try format.kind_codes.decode(try self.reader.readU32());
104         const limit = frame.values.items.len;
105         if (tag == .scan) {
106             if (self.count > format.max_scan_depth) return error.ScanTooDeep;
107             var scan = Scan{ .length = try self.reader.readInt(i64) };
108             const inits = try self.readIdList(limit, program_mod.max_scan_carries);
109             @memcpy(scan.inits[0..inits.len], inits);
110             scan.init_count = inits.len;
111             return self.open("scan_body", scan);
112         }
113         self.dims.clearRetainingCapacity();
114         self.integers.clearRetainingCapacity();
115         const op = try self.readOperation(tag, limit);
116         try self.bind(frame, &op);
117     }
118 
119     fn close(self: *Decoder) !void {
120         std.debug.assert(self.count >= 2);
121         const body = &self.frames[self.count - 1];
122         std.debug.assert(body.remaining == 0);
123         const count = try self.readCount(@sizeOf(u32));
124         const outputs = try body.builder.arena.allocator().alloc(Id, count);
125         for (outputs) |*output| output.* = try self.readId(body.values.items.len);
126         const view = program_mod.Subgraph{
127             .values = body.builder.values.items,
128             .operations = body.builder.operations.items,
129             .parameters = body.builder.parameters_list.items,
130             .outputs = outputs,
131         };
132         const op = derivedOperation(.{ .scan = .{
133             .length = body.scan.length,
134             .inits = body.scan.inits[0..body.scan.init_count],
135             .body = &view,
136         } });
137         try self.bind(&self.frames[self.count - 2], &op);
138         self.count -= 1;
139         body.values.deinit(self.allocator);
140         body.builder.deinit();
141     }
142 
143     fn finish(self: *Decoder, root: *Frame) !Program {
144         std.debug.assert(self.count == 1);
145         std.debug.assert(root.remaining == 0);
146         const count = try self.readCount(@sizeOf(u32));
147         const outputs = try self.allocator.alloc(Value, count);
148         defer self.allocator.free(outputs);
149         for (outputs) |*output| {
150             output.* = root.values.items[(try self.readId(root.values.items.len)).index];
151         }
152         try self.reader.expectDone();
153         return root.builder.finish(outputs);
154     }
155 
156     fn bind(self: *Decoder, frame: *Frame, op: *const Operation) !void {
157         var buffer: [max_operands]Value = undefined;
158         const args = tensor.interpret.arguments(Value, op, frame.values.items, &buffer);
159         const value = try frame.builder.operation(op, args);
160         std.debug.assert(value.id.index == frame.values.items.len);
161         std.debug.assert(frame.builder.values.items.len == frame.values.items.len + 1);
162         try frame.values.append(self.allocator, value);
163     }
164 
165     fn readOperation(self: *Decoder, tag: Tag, limit: usize) !Operation {
166         switch (tag) {
167             .parameter => {
168                 const result = try self.readType();
169                 return typedOperation(result, .{ .parameter = .{ .index = 0 } });
170             },
171             .constant => {
172                 const result = try self.readType();
173                 const payload = try self.reader.readLengthPrefixedBytes();
174                 return typedOperation(result, .{ .constant = .{ .payload = payload } });
175             },
176             .iota => {
177                 const result = try self.readType();
178                 const axis = try self.reader.readInt(i64);
179                 return typedOperation(result, .{ .iota = .{ .axis = axis } });
180             },
181             .custom_call => {
182                 const result = try self.readType();
183                 const target = try self.reader.readLengthPrefixedBytes();
184                 const version = try self.reader.readU32();
185                 const operands = try self.readIdList(limit, program_mod.max_custom_call_operands);
186                 return typedOperation(result, .{ .custom_call = .{
187                     .target = target,
188                     .version = version,
189                     .operands = operands,
190                 } });
191             },
192             .broadcast, .broadcast_in_dim, .reshape => return self.readShaped(tag, limit),
193             .unary, .binary, .compare, .select, .projection => {
194                 return self.readPointwise(tag, limit);
195             },
196             .transpose, .reduce, .gather, .scatter_add, .sparse_cross_entropy, .dot_general => {
197                 return self.readAxial(tag, limit);
198             },
199             .scan => unreachable,
200         }
201     }
202 
203     fn readShaped(self: *Decoder, tag: Tag, limit: usize) !Operation {
204         const result = Type{ .dtype = derived_result.dtype, .dims = try self.readDims() };
205         const input = try self.readId(limit);
206         return typedOperation(result, switch (tag) {
207             .broadcast => .{ .broadcast = .{ .input = input, .sizes = &.{} } },
208             .broadcast_in_dim => .{ .broadcast_in_dim = .{
209                 .input = input,
210                 .broadcast_dims = self.slice(try self.readIntegers()),
211             } },
212             .reshape => .{ .reshape = .{ .input = input, .new_shape = &.{} } },
213             else => unreachable,
214         });
215     }
216 
217     fn readPointwise(self: *Decoder, tag: Tag, limit: usize) !Operation {
218         switch (tag) {
219             .unary => {
220                 const op = try format.unary_codes.decode(try self.reader.readU32());
221                 const input = try self.readId(limit);
222                 return derivedOperation(.{ .unary = .{ .op = op, .input = input } });
223             },
224             .binary => {
225                 const op = try format.binary_codes.decode(try self.reader.readU32());
226                 const ids = try self.readIds(limit, 2);
227                 return derivedOperation(.{ .binary = .{ .op = op, .lhs = ids[0], .rhs = ids[1] } });
228             },
229             .compare => {
230                 const direction = try format.compare_codes.decode(try self.reader.readU32());
231                 const ids = try self.readIds(limit, 2);
232                 return derivedOperation(.{ .compare = .{
233                     .lhs = ids[0],
234                     .rhs = ids[1],
235                     .direction = direction,
236                 } });
237             },
238             .select => {
239                 const ids = try self.readIds(limit, 3);
240                 return derivedOperation(.{ .select = .{
241                     .pred = ids[0],
242                     .on_true = ids[1],
243                     .on_false = ids[2],
244                 } });
245             },
246             .projection => {
247                 const source = try self.readId(limit);
248                 const index = try self.reader.readU32();
249                 return derivedOperation(.{ .projection = .{ .source = source, .index = index } });
250             },
251             else => unreachable,
252         }
253     }
254 
255     fn readAxial(self: *Decoder, tag: Tag, limit: usize) !Operation {
256         switch (tag) {
257             .transpose => {
258                 const input = try self.readId(limit);
259                 const permutation = self.slice(try self.readIntegers());
260                 return derivedOperation(.{ .transpose = .{
261                     .input = input,
262                     .permutation = permutation,
263                 } });
264             },
265             .reduce => {
266                 const ids = try self.readIds(limit, 2);
267                 const reducer = try format.reducer_codes.decode(try self.reader.readU32());
268                 const dimensions = self.slice(try self.readIntegers());
269                 return derivedOperation(.{ .reduce = .{
270                     .input = ids[0],
271                     .init = ids[1],
272                     .reducer = reducer,
273                     .dimensions = dimensions,
274                 } });
275             },
276             .gather => {
277                 const ids = try self.readIds(limit, 2);
278                 const axis = try self.reader.readInt(i64);
279                 return derivedOperation(.{ .gather = .{
280                     .input = ids[0],
281                     .indices = ids[1],
282                     .axis = axis,
283                 } });
284             },
285             .scatter_add => {
286                 const ids = try self.readIds(limit, 3);
287                 const axis = try self.reader.readInt(i64);
288                 return derivedOperation(.{ .scatter_add = .{
289                     .input = ids[0],
290                     .indices = ids[1],
291                     .updates = ids[2],
292                     .axis = axis,
293                 } });
294             },
295             .sparse_cross_entropy, .dot_general => return self.readContraction(tag, limit),
296             else => unreachable,
297         }
298     }
299 
300     fn readContraction(self: *Decoder, tag: Tag, limit: usize) !Operation {
301         const ids = try self.readIds(limit, 2);
302         if (tag == .sparse_cross_entropy) {
303             const axis = try self.reader.readInt(i64);
304             return derivedOperation(.{ .sparse_cross_entropy = .{
305                 .logits = ids[0],
306                 .targets = ids[1],
307                 .axis = axis,
308             } });
309         }
310         std.debug.assert(tag == .dot_general);
311         var spans: [4]Span = undefined;
312         for (&spans) |*span| span.* = try self.readIntegers();
313         return derivedOperation(.{ .dot_general = .{
314             .lhs = ids[0],
315             .rhs = ids[1],
316             .lhs_contract = self.slice(spans[0]),
317             .rhs_contract = self.slice(spans[1]),
318             .lhs_batch = self.slice(spans[2]),
319             .rhs_batch = self.slice(spans[3]),
320         } });
321     }
322 
323     fn readCount(self: *Decoder, min_item_bytes: usize) !u32 {
324         std.debug.assert(min_item_bytes != 0);
325         const count = try self.reader.readU32();
326         if (count > self.reader.remaining() / min_item_bytes) return error.InvalidArtifact;
327         return count;
328     }
329 
330     fn readId(self: *Decoder, limit: usize) !Id {
331         const index = try self.reader.readU32();
332         if (index >= limit) return error.InvalidReference;
333         return .{ .index = index };
334     }
335 
336     fn readIds(self: *Decoder, limit: usize, comptime count: usize) ![count]Id {
337         var ids: [count]Id = undefined;
338         for (&ids) |*id| id.* = try self.readId(limit);
339         return ids;
340     }
341 
342     fn readIdList(self: *Decoder, limit: usize, max: usize) ![]const Id {
343         std.debug.assert(max <= self.operands.len);
344         const count = try self.readCount(@sizeOf(u32));
345         if (count > max) return error.TooManyOperands;
346         for (self.operands[0..count]) |*id| id.* = try self.readId(limit);
347         return self.operands[0..count];
348     }
349 
350     fn readIntegers(self: *Decoder) !Span {
351         const count = try self.readCount(@sizeOf(i64));
352         const start = self.integers.items.len;
353         try self.integers.ensureUnusedCapacity(self.allocator, count);
354         for (0..count) |_| self.integers.appendAssumeCapacity(try self.reader.readInt(i64));
355         return .{ .start = start, .end = self.integers.items.len };
356     }
357 
358     fn slice(self: *const Decoder, span: Span) []const i64 {
359         std.debug.assert(span.start <= span.end);
360         std.debug.assert(span.end <= self.integers.items.len);
361         return self.integers.items[span.start..span.end];
362     }
363 
364     fn readDims(self: *Decoder) ![]const Dim {
365         const rank = try self.readCount(min_dim_bytes);
366         self.dims.clearRetainingCapacity();
367         try self.dims.ensureTotalCapacity(self.allocator, rank);
368         for (0..rank) |_| {
369             const name = try self.reader.readLengthPrefixedBytes();
370             const extent = try self.reader.readInt(i64);
371             self.dims.appendAssumeCapacity(.{ .name = name, .extent = extent });
372         }
373         return self.dims.items;
374     }
375 
376     fn readType(self: *Decoder) !Type {
377         const dtype = try format.dtype_codes.decode(try self.reader.readU32());
378         return .{ .dtype = dtype, .dims = try self.readDims() };
379     }
380 };
381 
382 fn typedOperation(result: Type, kind: Kind) Operation {
383     return .{ .id = program_mod.synthetic_id, .result = result, .kind = kind };
384 }
385 
386 fn derivedOperation(kind: Kind) Operation {
387     return .{ .id = program_mod.synthetic_id, .result = derived_result, .kind = kind };
388 }