[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());