authorgravatar for john.schmidt.h@gmail.comJohn Schmidt <john.schmidt.h@gmail.com> 2022-03-15 23:25:38+01:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2022-03-16 20:11:05-07:00
logc8ed813097ebb679e858a7764673f6236e638ea4
tree2f0c506bbebe70a70e9678110f87d43cc204147f
parent312536540baf26728a56304811f63f01a7414b7a

Implement `@mulAdd` for vectors


3 files changed, 254 insertions(+), 58 deletions(-)

src/Sema.zig+76-48
...@@ -14499,19 +14499,24 @@ fn zirMulAdd(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air....@@ -14499,19 +14499,24 @@ fn zirMulAdd(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air.
1449914499
14500 const target = sema.mod.getTarget();14500 const target = sema.mod.getTarget();
1450114501
14502 const maybe_mulend1 = try sema.resolveMaybeUndefVal(block, mulend1_src, mulend1);
14503 const maybe_mulend2 = try sema.resolveMaybeUndefVal(block, mulend2_src, mulend2);
14504 const maybe_addend = try sema.resolveMaybeUndefVal(block, addend_src, addend);
14505
14502 switch (ty.zigTypeTag()) {14506 switch (ty.zigTypeTag()) {
14503 .ComptimeFloat, .Float => {14507 .ComptimeFloat, .Float, .Vector => {},
14504 const maybe_mulend1 = try sema.resolveMaybeUndefVal(block, mulend1_src, mulend1);14508 else => return sema.fail(block, src, "expected vector of floats or float type, found '{}'", .{ty}),
14505 const maybe_mulend2 = try sema.resolveMaybeUndefVal(block, mulend2_src, mulend2);14509 }
14506 const maybe_addend = try sema.resolveMaybeUndefVal(block, addend_src, addend);
1450714510
14508 const runtime_src = if (maybe_mulend1) |mulend1_val| rs: {14511 const runtime_src = if (maybe_mulend1) |mulend1_val| rs: {
14509 if (maybe_mulend2) |mulend2_val| {14512 if (maybe_mulend2) |mulend2_val| {
14510 if (mulend2_val.isUndef()) return sema.addConstUndef(ty);14513 if (mulend2_val.isUndef()) return sema.addConstUndef(ty);
1451114514
14512 if (maybe_addend) |addend_val| {14515 if (maybe_addend) |addend_val| {
14513 if (addend_val.isUndef()) return sema.addConstUndef(ty);14516 if (addend_val.isUndef()) return sema.addConstUndef(ty);
1451414517
14518 switch (ty.zigTypeTag()) {
14519 .ComptimeFloat, .Float => {
14515 const result_val = try Value.mulAdd(14520 const result_val = try Value.mulAdd(
14516 ty,14521 ty,
14517 mulend1_val,14522 mulend1_val,
...@@ -14521,47 +14526,70 @@ fn zirMulAdd(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air....@@ -14521,47 +14526,70 @@ fn zirMulAdd(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air.
14521 target,14526 target,
14522 );14527 );
14523 return sema.addConstant(ty, result_val);14528 return sema.addConstant(ty, result_val);
14524 } else {14529 },
14525 break :rs addend_src;14530 .Vector => {
14526 }14531 const scalar_ty = ty.scalarType();
14527 } else {14532 switch (scalar_ty.zigTypeTag()) {
14528 if (maybe_addend) |addend_val| {14533 .ComptimeFloat, .Float => {},
14529 if (addend_val.isUndef()) return sema.addConstUndef(ty);14534 else => return sema.fail(block, src, "expected vector of floats, found vector of '{}'", .{scalar_ty}),
14530 }14535 }
14531 break :rs mulend2_src;
14532 }
14533 } else rs: {
14534 if (maybe_mulend2) |mulend2_val| {
14535 if (mulend2_val.isUndef()) return sema.addConstUndef(ty);
14536 }
14537 if (maybe_addend) |addend_val| {
14538 if (addend_val.isUndef()) return sema.addConstUndef(ty);
14539 }
14540 break :rs mulend1_src;
14541 };
1454214536
14543 try sema.requireRuntimeBlock(block, runtime_src);14537 const vec_len = ty.vectorLen();
14544 return block.addInst(.{14538 const result_ty = try Type.vector(sema.arena, vec_len, scalar_ty);
14545 .tag = .mul_add,14539 var mulend1_buf: Value.ElemValueBuffer = undefined;
14546 .data = .{ .pl_op = .{14540 var mulend2_buf: Value.ElemValueBuffer = undefined;
14547 .operand = addend,14541 var addend_buf: Value.ElemValueBuffer = undefined;
14548 .payload = try sema.addExtra(Air.Bin{14542 const elems = try sema.arena.alloc(Value, vec_len);
14549 .lhs = mulend1,14543 for (elems) |*elem, i| {
14550 .rhs = mulend2,14544 const mulend1_elem_val = mulend1_val.elemValueBuffer(i, &mulend1_buf);
14551 }),14545 const mulend2_elem_val = mulend2_val.elemValueBuffer(i, &mulend2_buf);
14552 } },14546 const addend_elem_val = addend_val.elemValueBuffer(i, &addend_buf);
14553 });14547 elem.* = try Value.mulAdd(
14554 },14548 scalar_ty,
14555 .Vector => {14549 mulend1_elem_val,
14556 const scalar_ty = ty.scalarType();14550 mulend2_elem_val,
14557 switch (scalar_ty.zigTypeTag()) {14551 addend_elem_val,
14558 .ComptimeFloat, .Float => {},14552 sema.arena,
14559 else => return sema.fail(block, src, "expected vector of floats or float type, found '{}'", .{scalar_ty}),14553 target,
14554 );
14555 }
14556 return sema.addConstant(
14557 result_ty,
14558 try Value.Tag.aggregate.create(sema.arena, elems),
14559 );
14560 },
14561 else => unreachable,
14562 }
14563 } else {
14564 break :rs addend_src;
14560 }14565 }
14561 return sema.fail(block, src, "TODO: implement @mulAdd for vectors", .{});14566 } else {
14562 },14567 if (maybe_addend) |addend_val| {
14563 else => return sema.fail(block, src, "expected vector of floats or float type, found '{}'", .{ty}),14568 if (addend_val.isUndef()) return sema.addConstUndef(ty);
14564 }14569 }
14570 break :rs mulend2_src;
14571 }
14572 } else rs: {
14573 if (maybe_mulend2) |mulend2_val| {
14574 if (mulend2_val.isUndef()) return sema.addConstUndef(ty);
14575 }
14576 if (maybe_addend) |addend_val| {
14577 if (addend_val.isUndef()) return sema.addConstUndef(ty);
14578 }
14579 break :rs mulend1_src;
14580 };
14581
14582 try sema.requireRuntimeBlock(block, runtime_src);
14583 return block.addInst(.{
14584 .tag = .mul_add,
14585 .data = .{ .pl_op = .{
14586 .operand = addend,
14587 .payload = try sema.addExtra(Air.Bin{
14588 .lhs = mulend1,
14589 .rhs = mulend2,
14590 }),
14591 } },
14592 });
14565}14593}
1456614594
14567fn zirBuiltinCall(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air.Inst.Ref {14595fn zirBuiltinCall(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air.Inst.Ref {
src/codegen/llvm.zig+45-10
...@@ -5166,7 +5166,13 @@ pub const FuncGen = struct {...@@ -5166,7 +5166,13 @@ pub const FuncGen = struct {
5166 intrinsic,5166 intrinsic,
5167 libc: [*:0]const u8,5167 libc: [*:0]const u8,
5168 };5168 };
5169 const strat: Strat = switch (ty.floatBits(target)) {5169
5170 const scalar_ty = if (ty.zigTypeTag() == .Vector)
5171 ty.elemType()
5172 else
5173 ty;
5174
5175 const strat: Strat = switch (scalar_ty.floatBits(target)) {
5170 16, 32, 64 => Strat.intrinsic,5176 16, 32, 64 => Strat.intrinsic,
5171 80 => if (CType.longdouble.sizeInBits(target) == 80) Strat{ .intrinsic = {} } else Strat{ .libc = "__fmax" },5177 80 => if (CType.longdouble.sizeInBits(target) == 80) Strat{ .intrinsic = {} } else Strat{ .libc = "__fmax" },
5172 // LLVM always lowers the fma builtin for f128 to fmal, which is for `long double`.5178 // LLVM always lowers the fma builtin for f128 to fmal, which is for `long double`.
...@@ -5175,17 +5181,46 @@ pub const FuncGen = struct {...@@ -5175,17 +5181,46 @@ pub const FuncGen = struct {
5175 else => unreachable,5181 else => unreachable,
5176 };5182 };
51775183
5178 const llvm_fn = switch (strat) {5184 switch (strat) {
5179 .intrinsic => self.getIntrinsic("llvm.fma", &.{llvm_ty}),5185 .intrinsic => {
5180 .libc => |fn_name| self.dg.object.llvm_module.getNamedFunction(fn_name) orelse b: {5186 const llvm_fn = self.getIntrinsic("llvm.fma", &.{llvm_ty});
5181 const param_types = [_]*const llvm.Type{ llvm_ty, llvm_ty, llvm_ty };5187 const params = [_]*const llvm.Value{ mulend1, mulend2, addend };
5182 const fn_type = llvm.functionType(llvm_ty, &param_types, param_types.len, .False);5188 return self.builder.buildCall(llvm_fn, &params, params.len, .C, .Auto, "");
5183 break :b self.dg.object.llvm_module.addFunction(fn_name, fn_type);
5184 },5189 },
5185 };5190 .libc => |fn_name| {
5191 const scalar_llvm_ty = try self.dg.llvmType(scalar_ty);
5192 const llvm_fn = self.dg.object.llvm_module.getNamedFunction(fn_name) orelse b: {
5193 const param_types = [_]*const llvm.Type{ scalar_llvm_ty, scalar_llvm_ty, scalar_llvm_ty };
5194 const fn_type = llvm.functionType(scalar_llvm_ty, &param_types, param_types.len, .False);
5195 break :b self.dg.object.llvm_module.addFunction(fn_name, fn_type);
5196 };
5197
5198 if (ty.zigTypeTag() == .Vector) {
5199 const llvm_i32 = self.context.intType(32);
5200 const vector_llvm_ty = try self.dg.llvmType(ty);
5201
5202 var i: usize = 0;
5203 var vector = vector_llvm_ty.getUndef();
5204 while (i < ty.vectorLen()) : (i += 1) {
5205 const index_i32 = llvm_i32.constInt(i, .False);
5206
5207 const mulend1_elem = self.builder.buildExtractElement(mulend1, index_i32, "");
5208 const mulend2_elem = self.builder.buildExtractElement(mulend2, index_i32, "");
5209 const addend_elem = self.builder.buildExtractElement(addend, index_i32, "");
51865210
5187 const params = [_]*const llvm.Value{ mulend1, mulend2, addend };5211 const params = [_]*const llvm.Value{ mulend1_elem, mulend2_elem, addend_elem };
5188 return self.builder.buildCall(llvm_fn, &params, params.len, .C, .Auto, "");5212 const mul_add = self.builder.buildCall(llvm_fn, &params, params.len, .C, .Auto, "");
5213
5214 vector = self.builder.buildInsertElement(vector, mul_add, index_i32, "");
5215 }
5216
5217 return vector;
5218 } else {
5219 const params = [_]*const llvm.Value{ mulend1, mulend2, addend };
5220 return self.builder.buildCall(llvm_fn, &params, params.len, .C, .Auto, "");
5221 }
5222 },
5223 }
5189 }5224 }
51905225
5191 fn airShlWithOverflow(self: *FuncGen, inst: Air.Inst.Index) !?*const llvm.Value {5226 fn airShlWithOverflow(self: *FuncGen, inst: Air.Inst.Index) !?*const llvm.Value {
test/behavior/muladd.zig+133
...@@ -78,3 +78,136 @@ fn testMulAdd128() !void {...@@ -78,3 +78,136 @@ fn testMulAdd128() !void {
78 var c: f128 = 6.25;78 var c: f128 = 6.25;
79 try expect(@mulAdd(f128, a, b, c) == 20);79 try expect(@mulAdd(f128, a, b, c) == 20);
80}80}
81
82fn vector16() !void {
83 var a = @Vector(4, f16){ 5.5, 5.5, 5.5, 5.5 };
84 var b = @Vector(4, f16){ 2.5, 2.5, 2.5, 2.5 };
85 var c = @Vector(4, f16){ 6.25, 6.25, 6.25, 6.25 };
86 var x = @mulAdd(@Vector(4, f16), a, b, c);
87
88 // TODO use `expectEqual` instead once stage2 supports it
89 // var expected = @Vector(4, f16){ 20, 20, 20, 20 };
90 // try expectEqual(expected, x);
91
92 try expect(x[0] == 20);
93 try expect(x[1] == 20);
94 try expect(x[2] == 20);
95 try expect(x[3] == 20);
96}
97
98test "vector f16" {
99 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
100 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
101 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
102 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
103 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
104 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
105
106 comptime try vector16();
107 try vector16();
108}
109
110fn vector32() !void {
111 var a = @Vector(4, f32){ 5.5, 5.5, 5.5, 5.5 };
112 var b = @Vector(4, f32){ 2.5, 2.5, 2.5, 2.5 };
113 var c = @Vector(4, f32){ 6.25, 6.25, 6.25, 6.25 };
114 var x = @mulAdd(@Vector(4, f32), a, b, c);
115
116 // TODO use `expectEqual` instead once stage2 supports it
117 // var expected = @Vector(4, f32){ 20, 20, 20, 20 };
118 // try expectEqual(expected, x);
119
120 try expect(x[0] == 20);
121 try expect(x[1] == 20);
122 try expect(x[2] == 20);
123 try expect(x[3] == 20);
124}
125
126test "vector f32" {
127 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
128 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
129 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
130 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
131 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
132 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
133
134 comptime try vector32();
135 try vector32();
136}
137
138fn vector64() !void {
139 var a = @Vector(4, f64){ 5.5, 5.5, 5.5, 5.5 };
140 var b = @Vector(4, f64){ 2.5, 2.5, 2.5, 2.5 };
141 var c = @Vector(4, f64){ 6.25, 6.25, 6.25, 6.25 };
142 var x = @mulAdd(@Vector(4, f64), a, b, c);
143
144 // TODO use `expectEqual` instead once stage2 supports it
145 // var expected = @Vector(4, f64){ 20, 20, 20, 20 };
146 // try expectEqual(expected, x);
147
148 try expect(x[0] == 20);
149 try expect(x[1] == 20);
150 try expect(x[2] == 20);
151 try expect(x[3] == 20);
152}
153
154test "vector f64" {
155 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
156 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
157 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
158 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
159 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
160 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
161
162 comptime try vector64();
163 try vector64();
164}
165
166fn vector80() !void {
167 var a = @Vector(4, f80){ 5.5, 5.5, 5.5, 5.5 };
168 var b = @Vector(4, f80){ 2.5, 2.5, 2.5, 2.5 };
169 var c = @Vector(4, f80){ 6.25, 6.25, 6.25, 6.25 };
170 var x = @mulAdd(@Vector(4, f80), a, b, c);
171 try expect(x[0] == 20);
172 try expect(x[1] == 20);
173 try expect(x[2] == 20);
174 try expect(x[3] == 20);
175}
176
177test "vector f80" {
178 if (true) {
179 // https://github.com/ziglang/zig/issues/11030
180 return error.SkipZigTest;
181 }
182
183 comptime try vector80();
184 try vector80();
185}
186
187fn vector128() !void {
188 var a = @Vector(4, f128){ 5.5, 5.5, 5.5, 5.5 };
189 var b = @Vector(4, f128){ 2.5, 2.5, 2.5, 2.5 };
190 var c = @Vector(4, f128){ 6.25, 6.25, 6.25, 6.25 };
191 var x = @mulAdd(@Vector(4, f128), a, b, c);
192
193 // TODO use `expectEqual` instead once stage2 supports it
194 // var expected = @Vector(4, f128){ 20, 20, 20, 20 };
195 // try expectEqual(expected, x);
196
197 try expect(x[0] == 20);
198 try expect(x[1] == 20);
199 try expect(x[2] == 20);
200 try expect(x[3] == 20);
201}
202
203test "vector f128" {
204 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
205 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
206 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
207 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
208 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
209 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
210
211 comptime try vector128();
212 try vector128();
213}