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.
1449914499
1450014500 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
1450214506 switch (ty.zigTypeTag()) {
14503 .ComptimeFloat, .Float => {
14504 const maybe_mulend1 = try sema.resolveMaybeUndefVal(block, mulend1_src, mulend1);
14505 const maybe_mulend2 = try sema.resolveMaybeUndefVal(block, mulend2_src, mulend2);
14506 const maybe_addend = try sema.resolveMaybeUndefVal(block, addend_src, addend);
14507 .ComptimeFloat, .Float, .Vector => {},
14508 else => return sema.fail(block, src, "expected vector of floats or float type, found '{}'", .{ty}),
14509 }
1450714510
14508 const runtime_src = if (maybe_mulend1) |mulend1_val| rs: {
14509 if (maybe_mulend2) |mulend2_val| {
14510 if (mulend2_val.isUndef()) return sema.addConstUndef(ty);
14511 const runtime_src = if (maybe_mulend1) |mulend1_val| rs: {
14512 if (maybe_mulend2) |mulend2_val| {
14513 if (mulend2_val.isUndef()) return sema.addConstUndef(ty);
1451114514
14512 if (maybe_addend) |addend_val| {
14513 if (addend_val.isUndef()) return sema.addConstUndef(ty);
14515 if (maybe_addend) |addend_val| {
14516 if (addend_val.isUndef()) return sema.addConstUndef(ty);
1451414517
14518 switch (ty.zigTypeTag()) {
14519 .ComptimeFloat, .Float => {
1451514520 const result_val = try Value.mulAdd(
1451614521 ty,
1451714522 mulend1_val,
......@@ -14521,47 +14526,70 @@ fn zirMulAdd(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air.
1452114526 target,
1452214527 );
1452314528 return sema.addConstant(ty, result_val);
14524 } else {
14525 break :rs addend_src;
14526 }
14527 } else {
14528 if (maybe_addend) |addend_val| {
14529 if (addend_val.isUndef()) return sema.addConstUndef(ty);
14530 }
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 };
14529 },
14530 .Vector => {
14531 const scalar_ty = ty.scalarType();
14532 switch (scalar_ty.zigTypeTag()) {
14533 .ComptimeFloat, .Float => {},
14534 else => return sema.fail(block, src, "expected vector of floats, found vector of '{}'", .{scalar_ty}),
14535 }
1454214536
14543 try sema.requireRuntimeBlock(block, runtime_src);
14544 return block.addInst(.{
14545 .tag = .mul_add,
14546 .data = .{ .pl_op = .{
14547 .operand = addend,
14548 .payload = try sema.addExtra(Air.Bin{
14549 .lhs = mulend1,
14550 .rhs = mulend2,
14551 }),
14552 } },
14553 });
14554 },
14555 .Vector => {
14556 const scalar_ty = ty.scalarType();
14557 switch (scalar_ty.zigTypeTag()) {
14558 .ComptimeFloat, .Float => {},
14559 else => return sema.fail(block, src, "expected vector of floats or float type, found '{}'", .{scalar_ty}),
14537 const vec_len = ty.vectorLen();
14538 const result_ty = try Type.vector(sema.arena, vec_len, scalar_ty);
14539 var mulend1_buf: Value.ElemValueBuffer = undefined;
14540 var mulend2_buf: Value.ElemValueBuffer = undefined;
14541 var addend_buf: Value.ElemValueBuffer = undefined;
14542 const elems = try sema.arena.alloc(Value, vec_len);
14543 for (elems) |*elem, i| {
14544 const mulend1_elem_val = mulend1_val.elemValueBuffer(i, &mulend1_buf);
14545 const mulend2_elem_val = mulend2_val.elemValueBuffer(i, &mulend2_buf);
14546 const addend_elem_val = addend_val.elemValueBuffer(i, &addend_buf);
14547 elem.* = try Value.mulAdd(
14548 scalar_ty,
14549 mulend1_elem_val,
14550 mulend2_elem_val,
14551 addend_elem_val,
14552 sema.arena,
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;
1456014565 }
14561 return sema.fail(block, src, "TODO: implement @mulAdd for vectors", .{});
14562 },
14563 else => return sema.fail(block, src, "expected vector of floats or float type, found '{}'", .{ty}),
14564 }
14566 } else {
14567 if (maybe_addend) |addend_val| {
14568 if (addend_val.isUndef()) return sema.addConstUndef(ty);
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 });
1456514593}
1456614594
1456714595fn 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 {
51665166 intrinsic,
51675167 libc: [*:0]const u8,
51685168 };
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)) {
51705176 16, 32, 64 => Strat.intrinsic,
51715177 80 => if (CType.longdouble.sizeInBits(target) == 80) Strat{ .intrinsic = {} } else Strat{ .libc = "__fmax" },
51725178 // LLVM always lowers the fma builtin for f128 to fmal, which is for `long double`.
......@@ -5175,17 +5181,46 @@ pub const FuncGen = struct {
51755181 else => unreachable,
51765182 };
51775183
5178 const llvm_fn = switch (strat) {
5179 .intrinsic => self.getIntrinsic("llvm.fma", &.{llvm_ty}),
5180 .libc => |fn_name| self.dg.object.llvm_module.getNamedFunction(fn_name) orelse b: {
5181 const param_types = [_]*const llvm.Type{ llvm_ty, llvm_ty, llvm_ty };
5182 const fn_type = llvm.functionType(llvm_ty, &param_types, param_types.len, .False);
5183 break :b self.dg.object.llvm_module.addFunction(fn_name, fn_type);
5184 switch (strat) {
5185 .intrinsic => {
5186 const llvm_fn = self.getIntrinsic("llvm.fma", &.{llvm_ty});
5187 const params = [_]*const llvm.Value{ mulend1, mulend2, addend };
5188 return self.builder.buildCall(llvm_fn, &params, params.len, .C, .Auto, "");
51845189 },
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 };
5188 return self.builder.buildCall(llvm_fn, &params, params.len, .C, .Auto, "");
5211 const params = [_]*const llvm.Value{ mulend1_elem, mulend2_elem, addend_elem };
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 }
51895224 }
51905225
51915226 fn airShlWithOverflow(self: *FuncGen, inst: Air.Inst.Index) !?*const llvm.Value {
test/behavior/muladd.zig+117
......@@ -78,3 +78,120 @@ fn testMulAdd128() !void {
7878 var c: f128 = 6.25;
7979 try expect(@mulAdd(f128, a, b, c) == 20);
8080}
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}