lib/accy/src/tensor/trace/value.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const tensor = @import("../root.zig");
  2 const trace = @import("root.zig");
  3 
  4 const Builder = trace.Builder;
  5 const CompareDirection = trace.CompareDirection;
  6 const Dim = trace.Dim;
  7 const Id = trace.Id;
  8 const Reducer = trace.Reducer;
  9 const Type = trace.Type;
 10 const builder_mod = @import("builder.zig");
 11 const type_mod = tensor.types;
 12 
 13 pub const Value = struct {
 14     builder: *Builder,
 15     id: Id,
 16     ty: Type,
 17 
 18     pub fn add(self: Value, rhs: Value) !Value {
 19         return self.builder.alignedBinary(.add, self, rhs);
 20     }
 21 
 22     pub fn sub(self: Value, rhs: Value) !Value {
 23         return self.builder.alignedBinary(.sub, self, rhs);
 24     }
 25 
 26     pub fn mul(self: Value, rhs: Value) !Value {
 27         return self.builder.alignedBinary(.mul, self, rhs);
 28     }
 29 
 30     pub fn div(self: Value, rhs: Value) !Value {
 31         return self.builder.alignedBinary(.div, self, rhs);
 32     }
 33 
 34     pub fn pow(self: Value, rhs: Value) !Value {
 35         return self.builder.alignedBinary(.pow, self, rhs);
 36     }
 37 
 38     pub fn max(self: Value, other: anytype) !Value {
 39         if (comptime @TypeOf(other) == Value) {
 40             return self.builder.alignedBinary(.max, self, other);
 41         }
 42         return self.builder.reduceNamed(self, .max, valueAxes(other));
 43     }
 44 
 45     pub fn min(self: Value, other: anytype) !Value {
 46         if (comptime @TypeOf(other) == Value) {
 47             return self.builder.alignedBinary(.min, self, other);
 48         }
 49         return self.builder.reduceNamed(self, .min, valueAxes(other));
 50     }
 51 
 52     pub fn sum(self: Value, axes: anytype) !Value {
 53         return self.builder.reduceNamed(self, .sum, valueAxes(axes));
 54     }
 55 
 56     pub fn mean(self: Value, axes: anytype) !Value {
 57         return self.builder.meanNamed(self, valueAxes(axes));
 58     }
 59 
 60     pub fn reduce(self: Value, init: Value, reducer: Reducer, axes: anytype) !Value {
 61         return self.builder.reduceWith(self, init, reducer, valueAxes(axes));
 62     }
 63 
 64     pub fn gather(self: Value, indices: Value, comptime axis: anytype) !Value {
 65         return self.builder.gather(self, indices, axis);
 66     }
 67 
 68     pub fn scatterAdd(self: Value, indices: Value, updates: Value, comptime axis: anytype) !Value {
 69         return self.builder.scatterAdd(self, indices, updates, axis);
 70     }
 71 
 72     pub fn sparseCrossEntropyLoss(self: Value, targets: Value, comptime axis: anytype) !Value {
 73         return self.builder.sparseCrossEntropyLoss(self, targets, axis);
 74     }
 75 
 76     pub fn contract(self: Value, rhs: Value, axes: anytype) !Value {
 77         return self.builder.contract(self, rhs, valueAxes(axes));
 78     }
 79 
 80     pub fn rename(self: Value, comptime old_name: anytype, comptime new_name: anytype) !Value {
 81         return self.builder.renameAxis(self, comptime builder_mod.nameOf(old_name), comptime builder_mod.nameOf(new_name));
 82     }
 83 
 84     pub fn split(self: Value, comptime axis: anytype, parts_struct: anytype) !Value {
 85         var buffer: [type_mod.dimCount(@TypeOf(parts_struct))]Dim = undefined;
 86         type_mod.fillDims(parts_struct, &buffer);
 87         return self.builder.splitAxis(self, comptime builder_mod.nameOf(axis), &buffer);
 88     }
 89 
 90     pub fn merge(self: Value, axes: anytype, comptime merged_name: anytype) !Value {
 91         return self.builder.mergeAxes(self, valueAxes(axes), comptime builder_mod.nameOf(merged_name));
 92     }
 93 
 94     pub fn broadcast(self: Value, added_struct: anytype) !Value {
 95         var buffer: [type_mod.dimCount(@TypeOf(added_struct))]Dim = undefined;
 96         type_mod.fillDims(added_struct, &buffer);
 97         return self.builder.broadcastAxes(self, &buffer);
 98     }
 99 
100     pub fn isStructuralZero(self: Value) bool {
101         return self.builder.isStructuralZero(self.id);
102     }
103 
104     pub fn neg(self: Value) !Value {
105         return self.builder.unary(.neg, self);
106     }
107 
108     pub fn abs(self: Value) !Value {
109         return self.builder.unary(.abs, self);
110     }
111 
112     pub fn exp(self: Value) !Value {
113         return self.builder.unary(.exp, self);
114     }
115 
116     pub fn log(self: Value) !Value {
117         return self.builder.unary(.log, self);
118     }
119 
120     pub fn sqrt(self: Value) !Value {
121         return self.builder.unary(.sqrt, self);
122     }
123 
124     pub fn tanh(self: Value) !Value {
125         return self.builder.unary(.tanh, self);
126     }
127 
128     pub fn sin(self: Value) !Value {
129         return self.builder.unary(.sin, self);
130     }
131 
132     pub fn cos(self: Value) !Value {
133         return self.builder.unary(.cos, self);
134     }
135 
136     pub fn tan(self: Value) !Value {
137         return self.builder.unary(.tan, self);
138     }
139 
140     pub fn compare(self: Value, direction: CompareDirection, rhs: Value) !Value {
141         return self.builder.alignedCompare(direction, self, rhs);
142     }
143 
144     pub fn select(self: Value, on_true: Value, on_false: Value) !Value {
145         return self.builder.alignedSelect(self, on_true, on_false);
146     }
147 };
148 
149 fn valueAxes(axes: anytype) []const []const u8 {
150     const Axes = @TypeOf(axes);
151     if (comptime type_mod.isNameSlice(Axes)) {
152         return axes;
153     } else {
154         return comptime type_mod.axisNames(axes);
155     }
156 }