authorgravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2022-03-16 20:37:42-07:00
committergravatar for noreply@github.comGitHub <noreply@github.com> 2022-03-16 20:37:42-07:00
log1af51a083383be5bbb3249c24e00326aff935122
treed00ab6e1457220039868b041bb1ae891d594bd9d
parent312536540baf26728a56304811f63f01a7414b7a
parent80642b5984e5e7a59d380ced4c19e71c5e4cee32
signaturebadge-question-mark Signed by PGP key 4AEE18F83AFDEB23

Merge pull request #11196 from schmee/vector-muladd

Implement `@mulAdd` for vectors

3 files changed, 238 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+117
...@@ -78,3 +78,120 @@ fn testMulAdd128() !void {...@@ -78,3 +78,120 @@ 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 try expect(x[0] == 20);
89 try expect(x[1] == 20);
90 try expect(x[2] == 20);
91 try expect(x[3] == 20);
92}
93
94test "vector f16" {
95 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
96 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
97 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
98 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
99 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
100 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
101
102 comptime try vector16();
103 try vector16();
104}
105
106fn vector32() !void {
107 var a = @Vector(4, f32){ 5.5, 5.5, 5.5, 5.5 };
108 var b = @Vector(4, f32){ 2.5, 2.5, 2.5, 2.5 };
109 var c = @Vector(4, f32){ 6.25, 6.25, 6.25, 6.25 };
110 var x = @mulAdd(@Vector(4, f32), a, b, c);
111
112 try expect(x[0] == 20);
113 try expect(x[1] == 20);
114 try expect(x[2] == 20);
115 try expect(x[3] == 20);
116}
117
118test "vector f32" {
119 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
120 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
121 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
122 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
123 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
124 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
125
126 comptime try vector32();
127 try vector32();
128}
129
130fn vector64() !void {
131 var a = @Vector(4, f64){ 5.5, 5.5, 5.5, 5.5 };
132 var b = @Vector(4, f64){ 2.5, 2.5, 2.5, 2.5 };
133 var c = @Vector(4, f64){ 6.25, 6.25, 6.25, 6.25 };
134 var x = @mulAdd(@Vector(4, f64), a, b, c);
135
136 try expect(x[0] == 20);
137 try expect(x[1] == 20);
138 try expect(x[2] == 20);
139 try expect(x[3] == 20);
140}
141
142test "vector f64" {
143 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
144 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
145 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
146 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
147 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
148 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
149
150 comptime try vector64();
151 try vector64();
152}
153
154fn vector80() !void {
155 var a = @Vector(4, f80){ 5.5, 5.5, 5.5, 5.5 };
156 var b = @Vector(4, f80){ 2.5, 2.5, 2.5, 2.5 };
157 var c = @Vector(4, f80){ 6.25, 6.25, 6.25, 6.25 };
158 var x = @mulAdd(@Vector(4, f80), a, b, c);
159 try expect(x[0] == 20);
160 try expect(x[1] == 20);
161 try expect(x[2] == 20);
162 try expect(x[3] == 20);
163}
164
165test "vector f80" {
166 if (true) {
167 // https://github.com/ziglang/zig/issues/11030
168 return error.SkipZigTest;
169 }
170
171 comptime try vector80();
172 try vector80();
173}
174
175fn vector128() !void {
176 var a = @Vector(4, f128){ 5.5, 5.5, 5.5, 5.5 };
177 var b = @Vector(4, f128){ 2.5, 2.5, 2.5, 2.5 };
178 var c = @Vector(4, f128){ 6.25, 6.25, 6.25, 6.25 };
179 var x = @mulAdd(@Vector(4, f128), a, b, c);
180
181 try expect(x[0] == 20);
182 try expect(x[1] == 20);
183 try expect(x[2] == 20);
184 try expect(x[3] == 20);
185}
186
187test "vector f128" {
188 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
189 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
190 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
191 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
192 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
193 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
194
195 comptime try vector128();
196 try vector128();
197}