authorgravatar for robin@voetter.nlRobin Voetter <robin@voetter.nl> 2023-09-17 18:36:45+02:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2023-09-23 12:36:56-07:00
log98046b4c3c09d3d115b38a943fce328a2dc20dc4
treee42ec5c302a5bd32cbfcb08645a788afede54d48
parent6f55a689644fa85a6f650e6d33bcb882b1f507cd

spirv: air set_union_tag + improve load()/store()


2 files changed, 45 insertions(+), 60 deletions(-)

src/codegen/spirv.zig+45-41
......@@ -1662,13 +1662,11 @@ pub const DeclGen = struct {
16621662 return try self.convertToDirect(result_ty, result_id);
16631663 }
16641664
1665 fn load(self: *DeclGen, ptr_ty: Type, ptr_id: IdRef) !IdRef {
1666 const mod = self.module;
1667 const value_ty = ptr_ty.childType(mod);
1665 fn load(self: *DeclGen, value_ty: Type, ptr_id: IdRef, is_volatile: bool) !IdRef {
16681666 const indirect_value_ty_ref = try self.resolveType(value_ty, .indirect);
16691667 const result_id = self.spv.allocId();
16701668 const access = spec.MemoryAccess.Extended{
1671 .Volatile = ptr_ty.isVolatilePtr(mod),
1669 .Volatile = is_volatile,
16721670 };
16731671 try self.func.body.emit(self.spv.gpa, .OpLoad, .{
16741672 .id_result_type = self.typeId(indirect_value_ty_ref),
......@@ -1679,12 +1677,10 @@ pub const DeclGen = struct {
16791677 return try self.convertToDirect(value_ty, result_id);
16801678 }
16811679
1682 fn store(self: *DeclGen, ptr_ty: Type, ptr_id: IdRef, value_id: IdRef) !void {
1683 const mod = self.module;
1684 const value_ty = ptr_ty.childType(mod);
1680 fn store(self: *DeclGen, value_ty: Type, ptr_id: IdRef, value_id: IdRef, is_volatile: bool) !void {
16851681 const indirect_value_id = try self.convertToIndirect(value_ty, value_id);
16861682 const access = spec.MemoryAccess.Extended{
1687 .Volatile = ptr_ty.isVolatilePtr(mod),
1683 .Volatile = is_volatile,
16881684 };
16891685 try self.func.body.emit(self.spv.gpa, .OpStore, .{
16901686 .pointer = ptr_id,
......@@ -1754,6 +1750,7 @@ pub const DeclGen = struct {
17541750 .ptr_elem_ptr => try self.airPtrElemPtr(inst),
17551751 .ptr_elem_val => try self.airPtrElemVal(inst),
17561752
1753 .set_union_tag => return try self.airSetUnionTag(inst),
17571754 .get_union_tag => try self.airGetUnionTag(inst),
17581755 .struct_field_val => try self.airStructFieldVal(inst),
17591756
......@@ -2512,7 +2509,7 @@ pub const DeclGen = struct {
25122509
25132510 const slice_ptr = try self.extractField(ptr_ty, slice_id, 0);
25142511 const elem_ptr = try self.ptrAccessChain(ptr_ty_ref, slice_ptr, index_id, &.{});
2515 return try self.load(slice_ty, elem_ptr);
2512 return try self.load(slice_ty.childType(mod), elem_ptr, slice_ty.isVolatilePtr(mod));
25162513 }
25172514
25182515 fn ptrElemPtr(self: *DeclGen, ptr_ty: Type, ptr_id: IdRef, index_id: IdRef) !IdRef {
......@@ -2548,25 +2545,41 @@ pub const DeclGen = struct {
25482545 }
25492546
25502547 fn airPtrElemVal(self: *DeclGen, inst: Air.Inst.Index) !?IdRef {
2548 if (self.liveness.isUnused(inst)) return null;
2549
25512550 const mod = self.module;
25522551 const bin_op = self.air.instructions.items(.data)[inst].bin_op;
25532552 const ptr_ty = self.typeOf(bin_op.lhs);
2553 const elem_ty = self.typeOfIndex(inst);
25542554 const ptr_id = try self.resolve(bin_op.lhs);
25552555 const index_id = try self.resolve(bin_op.rhs);
2556
25572556 const elem_ptr_id = try self.ptrElemPtr(ptr_ty, ptr_id, index_id);
2557 return try self.load(elem_ty, elem_ptr_id, ptr_ty.isVolatilePtr(mod));
2558 }
2559
2560 fn airSetUnionTag(self: *DeclGen, inst: Air.Inst.Index) !void {
2561 const mod = self.module;
2562 const bin_op = self.air.instructions.items(.data)[inst].bin_op;
2563 const un_ptr_ty = self.typeOf(bin_op.lhs);
2564 const un_ty = un_ptr_ty.childType(mod);
2565 const layout = self.unionLayout(un_ty, null);
2566
2567 if (layout.tag_size == 0) return;
2568
2569 const tag_ty = un_ty.unionTagTypeSafety(mod).?;
2570 const tag_ty_ref = try self.resolveType(tag_ty, .indirect);
2571 const tag_ptr_ty_ref = try self.spv.ptrType(tag_ty_ref, spvStorageClass(un_ptr_ty.ptrAddressSpace(mod)));
25582572
2559 // If we have a pointer-to-array, construct an element pointer to use with load()
2560 // If we pass ptr_ty directly, it will attempt to load the entire array rather than
2561 // just an element.
2562 var elem_ptr_info = ptr_ty.ptrInfo(mod);
2563 elem_ptr_info.flags.size = .One;
2564 const elem_ptr_ty = try mod.intern_pool.get(mod.gpa, .{ .ptr_type = elem_ptr_info });
2573 const union_ptr_id = try self.resolve(bin_op.lhs);
2574 const new_tag_id = try self.resolve(bin_op.rhs);
25652575
2566 return try self.load(elem_ptr_ty.toType(), elem_ptr_id);
2576 const ptr_id = try self.accessChain(tag_ptr_ty_ref, union_ptr_id, &.{layout.tag_index});
2577 try self.store(tag_ty, ptr_id, new_tag_id, un_ptr_ty.isVolatilePtr(mod));
25672578 }
25682579
25692580 fn airGetUnionTag(self: *DeclGen, inst: Air.Inst.Index) !?IdRef {
2581 if (self.liveness.isUnused(inst)) return null;
2582
25702583 const ty_op = self.air.instructions.items(.data)[inst].ty_op;
25712584 const un_ty = self.typeOf(ty_op.operand);
25722585
......@@ -2588,25 +2601,25 @@ pub const DeclGen = struct {
25882601 const ty_pl = self.air.instructions.items(.data)[inst].ty_pl;
25892602 const struct_field = self.air.extraData(Air.StructField, ty_pl.payload).data;
25902603
2591 const container_ty = self.typeOf(struct_field.struct_operand);
2604 const object_ty = self.typeOf(struct_field.struct_operand);
25922605 const object_id = try self.resolve(struct_field.struct_operand);
25932606 const field_index = struct_field.field_index;
2594 const field_ty = container_ty.structFieldType(field_index, mod);
2607 const field_ty = object_ty.structFieldType(field_index, mod);
25952608
25962609 if (!field_ty.hasRuntimeBitsIgnoreComptime(mod)) return null;
25972610
2598 switch (container_ty.zigTypeTag(mod)) {
2599 .Struct => switch (container_ty.containerLayout(mod)) {
2611 switch (object_ty.zigTypeTag(mod)) {
2612 .Struct => switch (object_ty.containerLayout(mod)) {
26002613 .Packed => unreachable, // TODO
26012614 else => return try self.extractField(field_ty, object_id, field_index),
26022615 },
2603 .Union => switch (container_ty.containerLayout(mod)) {
2616 .Union => switch (object_ty.containerLayout(mod)) {
26042617 .Packed => unreachable, // TODO
26052618 else => {
26062619 // Store, pointer-cast, load
2607 const un_general_ty_ref = try self.resolveType(container_ty, .indirect);
2620 const un_general_ty_ref = try self.resolveType(object_ty, .indirect);
26082621 const un_general_ptr_ty_ref = try self.spv.ptrType(un_general_ty_ref, .Function);
2609 const un_active_ty_ref = try self.resolveUnionType(container_ty, field_index);
2622 const un_active_ty_ref = try self.resolveUnionType(object_ty, field_index);
26102623 const un_active_ptr_ty_ref = try self.spv.ptrType(un_active_ty_ref, .Function);
26112624 const field_ty_ref = try self.resolveType(field_ty, .indirect);
26122625 const field_ptr_ty_ref = try self.spv.ptrType(field_ty_ref, .Function);
......@@ -2617,31 +2630,20 @@ pub const DeclGen = struct {
26172630 .id_result = tmp_id,
26182631 .storage_class = .Function,
26192632 });
2620 try self.func.body.emit(self.spv.gpa, .OpStore, .{
2621 .pointer = tmp_id,
2622 .object = object_id,
2623 });
2633 try self.store(object_ty, tmp_id, object_id, false);
26242634 const casted_tmp_id = self.spv.allocId();
26252635 try self.func.body.emit(self.spv.gpa, .OpBitcast, .{
26262636 .id_result_type = self.typeId(un_active_ptr_ty_ref),
26272637 .id_result = casted_tmp_id,
26282638 .operand = tmp_id,
26292639 });
2630 const layout = self.unionLayout(container_ty, field_index);
2640 const layout = self.unionLayout(object_ty, field_index);
26312641 const field_ptr_id = try self.accessChain(field_ptr_ty_ref, casted_tmp_id, &.{layout.active_field_index});
2632 const result_id = self.spv.allocId();
2633 try self.func.body.emit(self.spv.gpa, .OpLoad, .{
2634 .id_result_type = self.typeId(field_ty_ref),
2635 .id_result = result_id,
2636 .pointer = field_ptr_id,
2637 });
2638 return try self.convertToDirect(field_ty, result_id);
2642 return try self.load(field_ty, field_ptr_id, false);
26392643 },
26402644 },
26412645 else => unreachable,
26422646 }
2643
2644 // return try self.extractField(field_ty, object_id, field_index);
26452647 }
26462648
26472649 fn structFieldPtr(
......@@ -2866,19 +2868,21 @@ pub const DeclGen = struct {
28662868 const mod = self.module;
28672869 const ty_op = self.air.instructions.items(.data)[inst].ty_op;
28682870 const ptr_ty = self.typeOf(ty_op.operand);
2871 const elem_ty = self.typeOfIndex(inst);
28692872 const operand = try self.resolve(ty_op.operand);
28702873 if (!ptr_ty.isVolatilePtr(mod) and self.liveness.isUnused(inst)) return null;
28712874
2872 return try self.load(ptr_ty, operand);
2875 return try self.load(elem_ty, operand, ptr_ty.isVolatilePtr(mod));
28732876 }
28742877
28752878 fn airStore(self: *DeclGen, inst: Air.Inst.Index) !void {
28762879 const bin_op = self.air.instructions.items(.data)[inst].bin_op;
28772880 const ptr_ty = self.typeOf(bin_op.lhs);
2881 const elem_ty = ptr_ty.childType(self.module);
28782882 const ptr = try self.resolve(bin_op.lhs);
28792883 const value = try self.resolve(bin_op.rhs);
28802884
2881 try self.store(ptr_ty, ptr, value);
2885 try self.store(elem_ty, ptr, value, ptr_ty.isVolatilePtr(self.module));
28822886 }
28832887
28842888 fn airLoop(self: *DeclGen, inst: Air.Inst.Index) !void {
......@@ -2922,7 +2926,7 @@ pub const DeclGen = struct {
29222926 }
29232927
29242928 const ptr = try self.resolve(un_op);
2925 const value = try self.load(ptr_ty, ptr);
2929 const value = try self.load(ret_ty, ptr, ptr_ty.isVolatilePtr(mod));
29262930 try self.func.body.emit(self.spv.gpa, .OpReturnValue, .{
29272931 .value = value,
29282932 });
test/behavior/union.zig-19
......@@ -29,7 +29,6 @@ test "init union with runtime value - floats" {
2929 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
3030 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
3131 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
32 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
3332
3433 var foo: FooWithFloats = undefined;
3534
......@@ -59,7 +58,6 @@ test "init union with runtime value" {
5958 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
6059 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
6160 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
62 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
6361
6462 var foo: Foo = undefined;
6563
......@@ -170,7 +168,6 @@ test "constant tagged union with payload" {
170168 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
171169 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
172170 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
173 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
174171
175172 var empty = TaggedUnionWithPayload{ .Empty = {} };
176173 var full = TaggedUnionWithPayload{ .Full = 13 };
......@@ -508,7 +505,6 @@ test "union initializer generates padding only if needed" {
508505test "runtime tag name with single field" {
509506 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
510507 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
511 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
512508
513509 const U = union(enum) {
514510 A: i32,
......@@ -585,7 +581,6 @@ test "tagged union as return value" {
585581 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
586582 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
587583 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
588 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
589584
590585 switch (returnAnInt(13)) {
591586 TaggedFoo.One => |value| try expect(value == 13),
......@@ -630,7 +625,6 @@ test "union(enum(u32)) with specified and unspecified tag values" {
630625 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
631626 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
632627 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
633 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
634628
635629 try comptime expect(Tag(Tag(MultipleChoice2)) == u32);
636630 try testEnumWithSpecifiedAndUnspecifiedTagValues(MultipleChoice2{ .C = 123 });
......@@ -668,7 +662,6 @@ test "switch on union with only 1 field" {
668662 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
669663 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
670664 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
671 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
672665
673666 var r: PartialInst = undefined;
674667 r = PartialInst.Compiled;
......@@ -697,7 +690,6 @@ const PartialInstWithPayload = union(enum) {
697690
698691test "union with only 1 field casted to its enum type which has enum value specified" {
699692 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
700 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
701693
702694 const Literal = union(enum) {
703695 Number: f64,
......@@ -782,7 +774,6 @@ test "return union init with void payload" {
782774 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
783775 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
784776 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
785 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
786777
787778 const S = struct {
788779 fn entry() !void {
......@@ -836,7 +827,6 @@ test "@unionInit can modify a union type" {
836827 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
837828 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
838829 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
839 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
840830
841831 const UnionInitEnum = union(enum) {
842832 Boolean: bool,
......@@ -860,7 +850,6 @@ test "@unionInit can modify a pointer value" {
860850 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
861851 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
862852 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
863 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
864853
865854 const UnionInitEnum = union(enum) {
866855 Boolean: bool,
......@@ -917,7 +906,6 @@ test "anonymous union literal syntax" {
917906 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
918907 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
919908 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
920 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
921909
922910 const S = struct {
923911 const Number = union {
......@@ -1041,7 +1029,6 @@ test "switching on non exhaustive union" {
10411029 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
10421030 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
10431031 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
1044 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
10451032
10461033 const S = struct {
10471034 const E = enum(u8) {
......@@ -1225,7 +1212,6 @@ test "union tag is set when initiated as a temporary value at runtime" {
12251212 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
12261213 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
12271214 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
1228 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
12291215
12301216 const U = union(enum) {
12311217 a,
......@@ -1263,7 +1249,6 @@ test "return an extern union from C calling convention" {
12631249 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
12641250 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
12651251 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
1266 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
12671252
12681253 const namespace = struct {
12691254 const S = extern struct {
......@@ -1294,7 +1279,6 @@ test "noreturn field in union" {
12941279 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
12951280 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
12961281 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
1297 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
12981282
12991283 const U = union(enum) {
13001284 a: u32,
......@@ -1475,7 +1459,6 @@ test "no dependency loop when function pointer in union returns the union" {
14751459test "union reassignment can use previous value" {
14761460 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
14771461 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
1478 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
14791462
14801463 const U = union {
14811464 a: u32,
......@@ -1527,7 +1510,6 @@ test "reinterpreting enum value inside packed union" {
15271510
15281511test "access the tag of a global tagged union" {
15291512 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
1530 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
15311513
15321514 const U = union(enum) {
15331515 a,
......@@ -1539,7 +1521,6 @@ test "access the tag of a global tagged union" {
15391521
15401522test "coerce enum literal to union in result loc" {
15411523 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
1542 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
15431524
15441525 const U = union(enum) {
15451526 a,