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 }