lib/accy/src/choir/memory.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir = @import("choir");
3 const contract = @import("contract.zig");
4 const dispatch = @import("dispatch.zig");
5 const semantic = @import("semantic.zig");
6 const tensor = @import("tensor.zig");
7
8 const ir = choir.ir;
9 const passes = choir.passes;
10
11 pub const product_name = "accy.memory";
12
13 pub const MemoryModule = @import("root.zig").publication.Module(.memory);
14
15 pub const MemoryJob = struct {
16 allocator: std.mem.Allocator,
17 dispatch_module: *dispatch.DispatchJob,
18 choir_module: *ir.Operation,
19 observed_fingerprint: u64,
20
21 pub fn init(
22 allocator: std.mem.Allocator,
23 dispatch_module: *dispatch.DispatchJob,
24 ) !*MemoryJob {
25 const module = try allocator.create(MemoryJob);
26 errdefer allocator.destroy(module);
27 module.* = .{
28 .allocator = allocator,
29 .dispatch_module = dispatch_module,
30 .choir_module = dispatch_module.choir_module,
31 .observed_fingerprint = 0,
32 };
33 try module.verify();
34 module.observed_fingerprint = try choir.operationFingerprint(allocator, module.choir_module);
35 return module;
36 }
37
38 pub fn context(self: *MemoryJob) *ir.Context {
39 return self.dispatch_module.context();
40 }
41
42 pub fn deinit(self: *MemoryJob) void {
43 self.dispatch_module.deinit();
44 const allocator = self.allocator;
45 allocator.destroy(self);
46 }
47
48 pub fn verify(self: *MemoryJob) !void {
49 try self.dispatch_module.verify();
50 }
51
52 pub fn fingerprint(self: *const MemoryJob) u64 {
53 return self.observed_fingerprint;
54 }
55
56 pub fn passContext(self: *MemoryJob) passes.PassContext {
57 return self.dispatch_module.passContext();
58 }
59 };
60
61 test "memory module owns dispatch product" {
62 const allocator = std.testing.allocator;
63
64 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
65 defer builder.deinit();
66
67 const ty = try builder.tensor(.f32, &.{4});
68 var function = try builder.beginFunction("memory_add", &.{ ty, ty }, &.{ty});
69 const sum = try function.add(function.parameter(0), function.parameter(1));
70 try function.return_(&.{sum});
71 try function.finish();
72
73 const semantic_module = try builder.finish();
74 var contract_module = try contract.ContractJob.init(allocator, semantic_module);
75 var contract_owned = true;
76 errdefer if (contract_owned) contract_module.deinit();
77
78 var analysis_cache = passes.AnalysisCache.init(allocator, null);
79 var cache_owned = true;
80 errdefer if (cache_owned) analysis_cache.deinit();
81
82 var tensor_module = try tensor.TensorJob.init(allocator, contract_module, analysis_cache);
83 contract_owned = false;
84 cache_owned = false;
85 var tensor_owned = true;
86 errdefer if (tensor_owned) tensor_module.deinit();
87
88 var dispatch_module = try dispatch.DispatchJob.init(allocator, tensor_module);
89 tensor_owned = false;
90 var dispatch_owned = true;
91 errdefer if (dispatch_owned) dispatch_module.deinit();
92
93 var memory_module = try MemoryJob.init(allocator, dispatch_module);
94 dispatch_owned = false;
95 defer memory_module.deinit();
96
97 try std.testing.expectEqualStrings(product_name, "accy.memory");
98 try memory_module.verify();
99 try std.testing.expectEqual(try choir.operationFingerprint(allocator, memory_module.choir_module), memory_module.fingerprint());
100 }