[hlsl] Add matrix storage store support This CL adds support for storing matrices in the DecomposeStorageAccess transform. Bug: 42251045 Change-Id: I67530a84972266a8e332d613f31d8ccdd38c587c Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/198335 Commit-Queue: dan sinclair <dsinclair@chromium.org> Reviewed-by: James Price <jrprice@google.com>
diff --git a/src/tint/lang/hlsl/writer/access_test.cc b/src/tint/lang/hlsl/writer/access_test.cc index 088e004..7265e23 100644 --- a/src/tint/lang/hlsl/writer/access_test.cc +++ b/src/tint/lang/hlsl/writer/access_test.cc
@@ -1449,7 +1449,7 @@ )"); } -TEST_F(HlslWriterTest, DISABLED_AccessStoreMatrixElement) { +TEST_F(HlslWriterTest, AccessStoreMatrixElement) { auto* var = b.Var<storage, mat4x4<f32>, core::Access::kReadWrite>("v"); var->SetBindingPoint(0, 0); @@ -1467,10 +1467,11 @@ void foo() { v.Store(24u, asuint(5.0f)); } + )"); } -TEST_F(HlslWriterTest, DISABLED_AccessStoreMatrixColumn) { +TEST_F(HlslWriterTest, AccessStoreMatrixColumn) { auto* var = b.Var<storage, mat4x4<f32>, core::Access::kReadWrite>("v"); var->SetBindingPoint(0, 0); @@ -1486,12 +1487,13 @@ EXPECT_EQ(output_.hlsl, R"( RWByteAddressBuffer v : register(u0); void foo() { - v2.Store4(16u, asuint((5.0f).xxxx)); + v.Store4(16u, asuint((5.0f).xxxx)); } + )"); } -TEST_F(HlslWriterTest, DISABLED_AccessStoreMatrix) { +TEST_F(HlslWriterTest, AccessStoreMatrix) { auto* var = b.Var<storage, mat4x4<f32>, core::Access::kReadWrite>("v"); var->SetBindingPoint(0, 0); @@ -1505,16 +1507,17 @@ ASSERT_TRUE(Generate()) << err_ << output_.hlsl; EXPECT_EQ(output_.hlsl, R"( RWByteAddressBuffer v : register(u0); -void v_1(uint offset, float4x4 value) { - v.Store4((offset + 0u), asuint(value[0u])); - v.Store4((offset + 16u), asuint(value[1u])); - v.Store4((offset + 32u), asuint(value[2u])); - v.Store4((offset + 48u), asuint(value[3u])); +void v_1(uint offset, float4x4 obj) { + v.Store4((offset + 0u), asuint(obj[0u])); + v.Store4((offset + 16u), asuint(obj[1u])); + v.Store4((offset + 32u), asuint(obj[2u])); + v.Store4((offset + 48u), asuint(obj[3u])); } void foo() { v_1(0u, float4x4((0.0f).xxxx, (0.0f).xxxx, (0.0f).xxxx, (0.0f).xxxx)); } + )"); }
diff --git a/src/tint/lang/hlsl/writer/raise/decompose_storage_access.cc b/src/tint/lang/hlsl/writer/raise/decompose_storage_access.cc index e76ac03..1140e12 100644 --- a/src/tint/lang/hlsl/writer/raise/decompose_storage_access.cc +++ b/src/tint/lang/hlsl/writer/raise/decompose_storage_access.cc
@@ -225,10 +225,10 @@ auto* fn = GetStoreFunctionFor(inst, var, s); b.Call(fn, offset, from); }, - // [&](const core::type::Matrix* m) { - // auto* fn = GetLoadFunctionFor(inst, var, m); - // return b.Call(fn, offset); - // }, / + [&](const core::type::Matrix* m) { + auto* fn = GetStoreFunctionFor(inst, var, m); + b.Call(fn, offset, from); + }, // [&](const core::type::Array* a) { // auto* fn = GetLoadFunctionFor(inst, var, a); // return b.Call(fn, offset); @@ -431,6 +431,30 @@ }); } + core::ir::Function* GetStoreFunctionFor(core::ir::Instruction* inst, + core::ir::Var* var, + const core::type::Matrix* mat) { + return var_and_type_to_store_fn_.GetOrAdd(VarTypePair{var, mat}, [&] { + auto* p = b.FunctionParam("offset", ty.u32()); + auto* obj = b.FunctionParam("obj", mat); + auto* fn = b.Function(ty.void_()); + fn->SetParams({p, obj}); + + b.Append(fn->Block(), [&] { + Vector<core::ir::Value*, 4> values; + for (size_t i = 0; i < mat->columns(); ++i) { + auto* from = b.Access(mat->ColumnType(), obj, u32(i)); + MakeStore(inst, var, from->Result(0), + b.Add<u32>(p, u32(i * mat->ColumnStride()))->Result(0)); + } + + b.Return(fn); + }); + + return fn; + }); + } + // Creates a load function for the given `var` and `array` combination. Essentially creates // a function similar to: //
diff --git a/src/tint/lang/hlsl/writer/raise/decompose_storage_access_test.cc b/src/tint/lang/hlsl/writer/raise/decompose_storage_access_test.cc index f0a14a7..1ab0097 100644 --- a/src/tint/lang/hlsl/writer/raise/decompose_storage_access_test.cc +++ b/src/tint/lang/hlsl/writer/raise/decompose_storage_access_test.cc
@@ -2594,26 +2594,26 @@ EXPECT_EQ(expect, str()); } -TEST_F(HlslWriterDecomposeStorageAccessTest, DISABLED_StoreMatrixElement) { - auto* var = b.Var<storage, mat4x4<f32>, core::Access::kReadWrite>("v"); +TEST_F(HlslWriterDecomposeStorageAccessTest, StoreMatrixElement) { + auto* var = b.Var<storage, mat2x3<f32>, core::Access::kReadWrite>("v"); var->SetBindingPoint(0, 0); b.ir.root_block->Append(var); auto* func = b.Function("foo", ty.void_(), core::ir::Function::PipelineStage::kFragment); b.Append(func->Block(), [&] { b.StoreVectorElement( - b.Access(ty.ptr<storage, vec4<f32>, core::Access::kReadWrite>(), var, 1_u), 2_u, 5_f); + b.Access(ty.ptr<storage, vec3<f32>, core::Access::kReadWrite>(), var, 1_u), 2_u, 5_f); b.Return(func); }); auto* src = R"( $B1: { # root - %v:ptr<storage, mat4x4<f32>, read_write> = var @binding_point(0, 0) + %v:ptr<storage, mat2x3<f32>, read_write> = var @binding_point(0, 0) } %foo = @fragment func():void { $B2: { - %3:ptr<storage, vec4<f32>, read_write> = access %v, 1u + %3:ptr<storage, vec3<f32>, read_write> = access %v, 1u store_vector_element %3, 2u, 5.0f ret } @@ -2622,32 +2622,43 @@ ASSERT_EQ(src, str()); auto* expect = R"( +$B1: { # root + %v:hlsl.byte_address_buffer<read_write> = var @binding_point(0, 0) +} + +%foo = @fragment func():void { + $B2: { + %3:u32 = bitcast 5.0f + %4:void = %v.Store 24u, %3 + ret + } +} )"; Run(DecomposeStorageAccess); EXPECT_EQ(expect, str()); } -TEST_F(HlslWriterDecomposeStorageAccessTest, DISABLED_StoreMatrixColumn) { - auto* var = b.Var<storage, mat4x4<f32>, core::Access::kReadWrite>("v"); +TEST_F(HlslWriterDecomposeStorageAccessTest, StoreMatrixColumn) { + auto* var = b.Var<storage, mat2x3<f32>, core::Access::kReadWrite>("v"); var->SetBindingPoint(0, 0); b.ir.root_block->Append(var); auto* func = b.Function("foo", ty.void_(), core::ir::Function::PipelineStage::kFragment); b.Append(func->Block(), [&] { - b.Store(b.Access(ty.ptr<storage, vec4<f32>, core::Access::kReadWrite>(), var, 1_u), - b.Splat<vec4<f32>>(5_f)); + b.Store(b.Access(ty.ptr<storage, vec3<f32>, core::Access::kReadWrite>(), var, 1_u), + b.Splat<vec3<f32>>(5_f)); b.Return(func); }); auto* src = R"( $B1: { # root - %v:ptr<storage, mat4x4<f32>, read_write> = var @binding_point(0, 0) + %v:ptr<storage, mat2x3<f32>, read_write> = var @binding_point(0, 0) } %foo = @fragment func():void { $B2: { - %3:ptr<storage, vec4<f32>, read_write> = access %v, 1u - store %3, vec4<f32>(5.0f) + %3:ptr<storage, vec3<f32>, read_write> = access %v, 1u + store %3, vec3<f32>(5.0f) ret } } @@ -2655,30 +2666,41 @@ ASSERT_EQ(src, str()); auto* expect = R"( +$B1: { # root + %v:hlsl.byte_address_buffer<read_write> = var @binding_point(0, 0) +} + +%foo = @fragment func():void { + $B2: { + %3:vec3<u32> = bitcast vec3<f32>(5.0f) + %4:void = %v.Store3 16u, %3 + ret + } +} )"; Run(DecomposeStorageAccess); EXPECT_EQ(expect, str()); } -TEST_F(HlslWriterDecomposeStorageAccessTest, DISABLED_StoreMatrix) { - auto* var = b.Var<storage, mat4x4<f32>, core::Access::kReadWrite>("v"); +TEST_F(HlslWriterDecomposeStorageAccessTest, StoreMatrix) { + auto* var = b.Var<storage, mat2x3<f32>, core::Access::kReadWrite>("v"); var->SetBindingPoint(0, 0); b.ir.root_block->Append(var); auto* func = b.Function("foo", ty.void_(), core::ir::Function::PipelineStage::kFragment); b.Append(func->Block(), [&] { - b.Store(var, b.Zero<mat4x4<f32>>()); + b.Store(var, b.Zero<mat2x3<f32>>()); b.Return(func); }); auto* src = R"( $B1: { # root - %v:ptr<storage, mat4x4<f32>, read_write> = var @binding_point(0, 0) + %v:ptr<storage, mat2x3<f32>, read_write> = var @binding_point(0, 0) } %foo = @fragment func():void { $B2: { - store %v, mat4x4<f32>(vec4<f32>(0.0f)) + store %v, mat2x3<f32>(vec3<f32>(0.0f)) ret } } @@ -2686,6 +2708,29 @@ ASSERT_EQ(src, str()); auto* expect = R"( +$B1: { # root + %v:hlsl.byte_address_buffer<read_write> = var @binding_point(0, 0) +} + +%foo = @fragment func():void { + $B2: { + %3:void = call %4, 0u, mat2x3<f32>(vec3<f32>(0.0f)) + ret + } +} +%4 = func(%offset:u32, %obj:mat2x3<f32>):void { + $B3: { + %7:vec3<f32> = access %obj, 0u + %8:u32 = add %offset, 0u + %9:vec3<u32> = bitcast %7 + %10:void = %v.Store3 %8, %9 + %11:vec3<f32> = access %obj, 1u + %12:u32 = add %offset, 16u + %13:vec3<u32> = bitcast %11 + %14:void = %v.Store3 %12, %13 + ret + } +} )"; Run(DecomposeStorageAccess); EXPECT_EQ(expect, str());