| author | |
| committer | |
| log | 1af51a083383be5bbb3249c24e00326aff935122 |
| tree | d00ab6e1457220039868b041bb1ae891d594bd9d |
| parent | 312536540baf26728a56304811f63f01a7414b7a |
| parent | 80642b5984e5e7a59d380ced4c19e71c5e4cee32 |
| signature |
Implement `@mulAdd` for vectors3 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. |
| 14499 | 14499 | ||
| 14500 | const target = sema.mod.getTarget(); | 14500 | const target = sema.mod.getTarget(); |
| 14501 | 14501 | ||
| 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); | ||
| 14507 | 14510 | ||
| 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); |
| 14511 | 14514 | ||
| 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); |
| 14514 | 14517 | ||
| 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 | }; | ||
| 14542 | 14536 | ||
| 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 | } |
| 14566 | 14594 | ||
| 14567 | fn zirBuiltinCall(sema: *Sema, block: *Block, inst: Zir.Inst.Index) CompileError!Air.Inst.Ref { | 14595 | fn 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 | }; |
| 5177 | 5183 | ||
| 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, ""); | ||
| 5186 | 5210 | ||
| 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 | } |
| 5190 | 5225 | ||
| 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 | |||
| 82 | fn 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 | |||
| 94 | test "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 | |||
| 106 | fn 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 | |||
| 118 | test "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 | |||
| 130 | fn 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 | |||
| 142 | test "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 | |||
| 154 | fn 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 | |||
| 165 | test "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 | |||
| 175 | fn 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 | |||
| 187 | test "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 | } |