Reland "[dawn] Switch indirect dispatch validation to use immediates"

This is a reland of commit 104068a1a6c0bacad1e9896f18a7c558f82f5444

New changes are the use of a designated initializer for Params and the
update to use the spanified APISetImmediates entrypoint.

Original change's description:
> [dawn] Switch indirect dispatch validation to use immediates
>
> Change the indirect dispatch validation shader to use immediates instead
> of a uniform buffer.
>
> Fix: 488346117
> Change-Id: Ib0804f399bd387bcd93b5540490b566d57ff88d4
> Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/294716
> Commit-Queue: Corentin Wallez <cwallez@chromium.org>
> Reviewed-by: Corentin Wallez <cwallez@chromium.org>

Change-Id: I9787cef71529e916f06cf65e2a02726dfc5b85b9
Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/298275
Reviewed-by: Shaobo Yan <shaoboyan@microsoft.com>
Reviewed-by: Loko Kung <lokokung@google.com>
Commit-Queue: Corentin Wallez <cwallez@chromium.org>
diff --git a/src/dawn/native/ComputePassEncoder.cpp b/src/dawn/native/ComputePassEncoder.cpp
index c299b84..d69d38b 100644
--- a/src/dawn/native/ComputePassEncoder.cpp
+++ b/src/dawn/native/ComputePassEncoder.cpp
@@ -52,6 +52,18 @@
 
 namespace {
 
+// Neither 'enableValidation' nor 'duplicateNumWorkgroups' can be declared as 'bool' as
+// currently in WGSL type 'bool' cannot be used in address space 'uniform' as 'it is
+// non-host-shareable'.
+struct IndirectDispatchParams {
+    uint32_t maxComputeWorkgroupsPerDimension;
+    uint32_t clientOffsetInU32;
+    uint32_t enableValidation;
+    uint32_t duplicateNumWorkgroups;
+    uint32_t linearIndexing;
+    uint32_t overflowValue;
+};
+
 ResultOrError<ComputePipelineBase*> GetOrCreateIndirectDispatchValidationPipeline(
     DeviceBase* device) {
     InternalPipelineStore* store = device->GetInternalPipelineStore();
@@ -60,13 +72,12 @@
         return store->dispatchIndirectValidationPipeline.Get();
     }
 
-    // TODO(https://crbug.com/dawn/488346117): Use immediates instead of uniform.
     // TODO(https://crbug.com/dawn/1108): Propagate validation feedback from this
     // shader in various failure modes.
     // Type 'bool' cannot be used in address space 'uniform' as it is non-host-shareable.
     Ref<ShaderModuleBase> shaderModule;
     DAWN_TRY_ASSIGN(shaderModule, utils::CreateShaderModule(device, DAWN_MULTILINE(
-        struct UniformParams {
+        struct Params {
             maxComputeWorkgroupsPerDimension: u32,
             clientOffsetInU32: u32,
             enableValidation: u32,
@@ -83,23 +94,23 @@
             data: array<u32>
         }
 
-        @group(0) @binding(0) var<uniform> uniformParams: UniformParams;
-        @group(0) @binding(1) var<storage, read_write> clientParams: IndirectParams;
-        @group(0) @binding(2) var<storage, read_write> validatedParams: ValidatedParams;
+        var<immediate> params: Params;
+        @group(0) @binding(0) var<storage, read_write> clientParams: IndirectParams;
+        @group(0) @binding(1) var<storage, read_write> validatedParams: ValidatedParams;
 
         @compute @workgroup_size(1, 1, 1)
         fn main() {
-            var workgroups = vec3u(clientParams.data[uniformParams.clientOffsetInU32 + 0],
-                                   clientParams.data[uniformParams.clientOffsetInU32 + 1],
-                                   clientParams.data[uniformParams.clientOffsetInU32 + 2]);
-            if (uniformParams.enableValidation > 0u) {
+            var workgroups = vec3u(clientParams.data[params.clientOffsetInU32 + 0],
+                                   clientParams.data[params.clientOffsetInU32 + 1],
+                                   clientParams.data[params.clientOffsetInU32 + 2]);
+            if (params.enableValidation > 0u) {
                 var invalid = false;
-                if (max(workgroups.x, max(workgroups.y, workgroups.z)) > uniformParams.maxComputeWorkgroupsPerDimension) {
+                if (max(workgroups.x, max(workgroups.y, workgroups.z)) > params.maxComputeWorkgroupsPerDimension) {
                     invalid = true;
-                } else if (uniformParams.linearIndexing > 0u) {
-                    invalid |= workgroups.x > (uniformParams.overflowValue / workgroups.y);
+                } else if (params.linearIndexing > 0u) {
+                    invalid |= workgroups.x > (params.overflowValue / workgroups.y);
                     let xy = workgroups.x * workgroups.y;
-                    invalid |= xy > (uniformParams.overflowValue / workgroups.z);
+                    invalid |= xy > (params.overflowValue / workgroups.z);
                 }
 
                 if (invalid) {
@@ -109,7 +120,7 @@
             validatedParams.data[0] = workgroups.x;
             validatedParams.data[1] = workgroups.y;
             validatedParams.data[2] = workgroups.z;
-            if (uniformParams.duplicateNumWorkgroups > 0u) {
+            if (params.duplicateNumWorkgroups > 0u) {
                 validatedParams.data[3] = workgroups.x;
                 validatedParams.data[4] = workgroups.y;
                 validatedParams.data[5] = workgroups.z;
@@ -122,14 +133,14 @@
                     utils::MakeBindGroupLayout(
                         device,
                         {
-                            {0, wgpu::ShaderStage::Compute, wgpu::BufferBindingType::Uniform},
-                            {1, wgpu::ShaderStage::Compute, kInternalStorageBufferBinding},
-                            {2, wgpu::ShaderStage::Compute, wgpu::BufferBindingType::Storage},
+                            {0, wgpu::ShaderStage::Compute, kInternalStorageBufferBinding},
+                            {1, wgpu::ShaderStage::Compute, wgpu::BufferBindingType::Storage},
                         },
                         /* allowInternalBinding */ true));
 
     Ref<PipelineLayoutBase> pipelineLayout;
-    DAWN_TRY_ASSIGN(pipelineLayout, utils::MakeBasicPipelineLayout(device, bindGroupLayout));
+    DAWN_TRY_ASSIGN(pipelineLayout, utils::MakeBasicPipelineLayout(device, bindGroupLayout,
+                                                                   sizeof(IndirectDispatchParams)));
 
     ComputePipelineDescriptor computePipelineDescriptor = {};
     computePipelineDescriptor.layout = pipelineLayout.Get();
@@ -353,34 +364,17 @@
     const uint64_t clientIndirectBindingSize =
         kDispatchIndirectSize + clientOffsetFromAlignedBoundary;
 
-    // Neither 'enableValidation' nor 'duplicateNumWorkgroups' can be declared as 'bool' as
-    // currently in WGSL type 'bool' cannot be used in address space 'uniform' as 'it is
-    // non-host-shareable'.
-    struct UniformParams {
-        uint32_t maxComputeWorkgroupsPerDimension;
-        uint32_t clientOffsetInU32;
-        uint32_t enableValidation;
-        uint32_t duplicateNumWorkgroups;
-        uint32_t linearIndexing;
-        uint32_t overflowValue;
+    // Set the immediate params.
+    IndirectDispatchParams params = {
+        .maxComputeWorkgroupsPerDimension = device->GetLimits().v1.maxComputeWorkgroupsPerDimension,
+        .clientOffsetInU32 =
+            static_cast<uint32_t>(clientOffsetFromAlignedBoundary / sizeof(uint32_t)),
+        .enableValidation = static_cast<uint32_t>(IsValidationEnabled()),
+        .duplicateNumWorkgroups = static_cast<uint32_t>(shouldDuplicateNumWorkgroups),
+        .linearIndexing = static_cast<uint32_t>(usesLinearIndexing),
+        .overflowValue = overflowValue,
     };
 
-    // Create a uniform buffer to hold parameters for the shader.
-    Ref<BufferBase> uniformBuffer;
-    {
-        UniformParams params = {};
-        params.maxComputeWorkgroupsPerDimension =
-            device->GetLimits().v1.maxComputeWorkgroupsPerDimension;
-        params.clientOffsetInU32 = clientOffsetFromAlignedBoundary / sizeof(uint32_t);
-        params.enableValidation = static_cast<uint32_t>(IsValidationEnabled());
-        params.duplicateNumWorkgroups = static_cast<uint32_t>(shouldDuplicateNumWorkgroups);
-        params.linearIndexing = static_cast<uint32_t>(usesLinearIndexing);
-        params.overflowValue = overflowValue;
-
-        DAWN_TRY_ASSIGN(uniformBuffer,
-                        utils::CreateBufferFromData(device, wgpu::BufferUsage::Uniform, {params}));
-    }
-
     // Reserve space in the scratch buffer to hold the validated indirect params.
     ScratchBuffer& scratchBuffer = store->scratchIndirectStorage;
     const uint64_t scratchBufferSize =
@@ -393,15 +387,15 @@
     DAWN_TRY_ASSIGN(validationBindGroup,
                     utils::MakeBindGroup(device, layout,
                                          {
-                                             {0, uniformBuffer},
-                                             {1, indirectBuffer, clientIndirectBindingOffset,
+                                             {0, indirectBuffer, clientIndirectBindingOffset,
                                               clientIndirectBindingSize},
-                                             {2, validatedIndirectBuffer, 0, scratchBufferSize},
+                                             {1, validatedIndirectBuffer, 0, scratchBufferSize},
                                          },
                                          UsageValidationMode::Internal));
 
     // Issue commands to validate the indirect buffer.
     APISetPipeline(validationPipeline.Get());
+    APISetImmediates(0, ByteSpanFromRef(params));
     APISetBindGroup(0, validationBindGroup.Get());
     APIDispatchWorkgroups(1);
 
diff --git a/src/dawn/native/utils/WGPUHelpers.cpp b/src/dawn/native/utils/WGPUHelpers.cpp
index 20381ac..037b460 100644
--- a/src/dawn/native/utils/WGPUHelpers.cpp
+++ b/src/dawn/native/utils/WGPUHelpers.cpp
@@ -79,9 +79,11 @@
 
 ResultOrError<Ref<PipelineLayoutBase>> MakeBasicPipelineLayout(
     DeviceBase* device,
-    const Ref<BindGroupLayoutBase>& bindGroupLayout) {
+    const Ref<BindGroupLayoutBase>& bindGroupLayout,
+    uint32_t immediateSize) {
     PipelineLayoutDescriptor descriptor;
     descriptor.bindGroupLayouts = SpanFromRef<BindGroupIndex>(bindGroupLayout.Get());
+    descriptor.immediateSize = immediateSize;
     return device->CreatePipelineLayout(&descriptor);
 }
 
diff --git a/src/dawn/native/utils/WGPUHelpers.h b/src/dawn/native/utils/WGPUHelpers.h
index b1616de..6861efe 100644
--- a/src/dawn/native/utils/WGPUHelpers.h
+++ b/src/dawn/native/utils/WGPUHelpers.h
@@ -74,7 +74,8 @@
 
 ResultOrError<Ref<PipelineLayoutBase>> MakeBasicPipelineLayout(
     DeviceBase* device,
-    const Ref<BindGroupLayoutBase>& bindGroupLayout);
+    const Ref<BindGroupLayoutBase>& bindGroupLayout,
+    uint32_t immediateSize = 0);
 
 // Helpers to make creating bind group layouts look nicer:
 //
diff --git a/src/dawn/tests/unittests/native/CommandBufferEncodingTests.cpp b/src/dawn/tests/unittests/native/CommandBufferEncodingTests.cpp
index 10666cd..febce62 100644
--- a/src/dawn/tests/unittests/native/CommandBufferEncodingTests.cpp
+++ b/src/dawn/tests/unittests/native/CommandBufferEncodingTests.cpp
@@ -215,6 +215,13 @@
         }
     };
 
+    auto ExpectSetValidationImmediates = [&](CommandIterator* commands) {
+        auto* cmd = commands->NextCommand<SetImmediatesCmd>();
+        commands->NextData<uint8_t>(cmd->size);
+        ASSERT_EQ(cmd->offset, 0u);
+        ASSERT_EQ(cmd->size, 6 * sizeof(uint32_t));
+    };
+
     auto ExpectSetValidationBindGroup = [&](CommandIterator* commands) {
         auto* cmd = commands->NextCommand<SetBindGroupCmd>();
         ASSERT_EQ(cmd->index, BindGroupIndex(0u));
@@ -241,6 +248,7 @@
 
             // Expect the validation.
             {Command::SetComputePipeline, ExpectSetValidationPipeline},
+            {Command::SetImmediates, ExpectSetValidationImmediates},
             {Command::SetBindGroup, ExpectSetValidationBindGroup},
             {Command::Dispatch, ExpectSetValidationDispatch},
 
@@ -254,6 +262,7 @@
 
             // Expect the validation.
             {Command::SetComputePipeline, ExpectSetValidationPipeline},
+            {Command::SetImmediates, ExpectSetValidationImmediates},
             {Command::SetBindGroup, ExpectSetValidationBindGroup},
             {Command::Dispatch, ExpectSetValidationDispatch},
 
@@ -272,6 +281,7 @@
 
             // Expect the validation.
             {Command::SetComputePipeline, ExpectSetValidationPipeline},
+            {Command::SetImmediates, ExpectSetValidationImmediates},
             {Command::SetBindGroup, ExpectSetValidationBindGroup},
             {Command::Dispatch, ExpectSetValidationDispatch},