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 }