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 }