Skip to content

Commit cf8b909

Browse files
committed
Arm backend: Add bfloat16 support to VGF backend.
- Add bf16 extension to default VgfCompileSpec - Handle bf16 in VGFSetup.sh - Needs bumping of Vulkan SDK to 1.4.350.0 to include VK_FORMAT_R16_SFLOAT_FPENCODING_BFLOAT16_ARM Initially tested with a single operator test of matmul. Signed-off-by: Erik Lundell <erik.lundell@arm.com> Change-Id: I74b0c15b5a4f9194c437e8e69d2349e9c282878b
1 parent 2e328aa commit cf8b909

6 files changed

Lines changed: 23 additions & 16 deletions

File tree

backends/arm/runtime/VGFSetup.cpp

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,7 @@ enum class FormatScalarKind {
6666
Uint,
6767
Sint,
6868
Float,
69+
BFloat,
6970
};
7071

7172
struct FormatInfo {
@@ -157,6 +158,7 @@ static uint32_t get_format_component_count(VkFormat format) {
157158
case VK_FORMAT_R16_UINT:
158159
case VK_FORMAT_R16_SINT:
159160
case VK_FORMAT_R16_SFLOAT:
161+
case VK_FORMAT_R16_SFLOAT_FPENCODING_BFLOAT16_ARM:
160162
case VK_FORMAT_R32_UINT:
161163
case VK_FORMAT_R32_SINT:
162164
case VK_FORMAT_R32_SFLOAT:
@@ -209,6 +211,9 @@ static bool get_format_info(VkFormat format, FormatInfo* info) {
209211
case VK_FORMAT_R16_SFLOAT:
210212
*info = FormatInfo{1, 2, FormatScalarKind::Float};
211213
return true;
214+
case VK_FORMAT_R16_SFLOAT_FPENCODING_BFLOAT16_ARM:
215+
*info = FormatInfo{1, 2, FormatScalarKind::BFloat};
216+
return true;
212217
case VK_FORMAT_R32_UINT:
213218
*info = FormatInfo{1, 4, FormatScalarKind::Uint};
214219
return true;
@@ -3615,6 +3620,7 @@ static uint32_t get_format_size(VkFormat format) {
36153620
case VK_FORMAT_R16_UINT:
36163621
case VK_FORMAT_R16_SINT:
36173622
case VK_FORMAT_R16_SFLOAT:
3623+
case VK_FORMAT_R16_SFLOAT_FPENCODING_BFLOAT16_ARM:
36183624
case VK_FORMAT_R8G8_UINT:
36193625
case VK_FORMAT_R8G8_SINT:
36203626
return 2;

backends/arm/scripts/vulkan_utils.sh

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -26,21 +26,21 @@ vulkan_sdk_arch="${ARCH}"
2626
# macOS and Linux x86_64 use the official LunarG SDK tarballs. Linux ARM64
2727
# uses a separately repackaged mirror of the same SDK version.
2828
if [[ "${os_name}" == "Darwin" ]]; then
29-
vulkan_sdk_version="1.4.341.1"
29+
vulkan_sdk_version="1.4.350.0"
3030
vulkan_sdk_arch="macOS"
3131
vulkan_sdk_url="https://sdk.lunarg.com/sdk/download/${vulkan_sdk_version}/mac/vulkansdk-macos-${vulkan_sdk_version}.zip"
32-
vulkan_sdk_sha256="632cbe96c8ed6ed00c6ce25e3a7738c466134f76586e1c51f1419410d7f9042e"
32+
vulkan_sdk_sha256="7acc181b8fd9b4781bf51ed086222ec95d22004b85b3d0a6683a7e48ca5a1679"
3333
elif [[ "${os_name}" == "Linux" ]] && [[ "${ARCH}" == "x86_64" ]]; then
34-
vulkan_sdk_version="1.4.341.1"
34+
vulkan_sdk_version="1.4.350.0"
3535
vulkan_sdk_url="https://sdk.lunarg.com/sdk/download/${vulkan_sdk_version}/linux/vulkansdk-linux-x86_64-${vulkan_sdk_version}.tar.xz"
36-
vulkan_sdk_sha256="3bf0f762afb6c79bc6a9d9fb5998745ccff928800a29619b501ed9de7fd9789b"
36+
vulkan_sdk_sha256="b65f068ab36263559da49d7cacd7e7b9df23824ca8b68ccc522a2b06f5725df2"
3737
elif [[ "${os_name}" == "Linux" ]] && ([[ "${ARCH}" == "aarch64" ]] || [[ "${ARCH}" == "arm64" ]]); then
38-
vulkan_sdk_version="1.4.341.1"
38+
vulkan_sdk_version="1.4.350.0"
3939
if [[ "${vulkan_sdk_arch}" == "arm64" ]]; then
4040
vulkan_sdk_arch="aarch64"
4141
fi
4242
vulkan_sdk_url="https://github.com/jakoch/vulkan-sdk-arm/releases/download/${vulkan_sdk_version}/vulkansdk-ubuntu-22.04-arm-${vulkan_sdk_version}.tar.xz"
43-
vulkan_sdk_sha256="345312aee2c835e128b30653278593f899a659a7ba287c571cafb22acb708b8f"
43+
vulkan_sdk_sha256="9e403d444219bb7c17e9231b580d704453e2afa30a1c2fdd568d1776dc68790b"
4444
else
4545
log_step "vulkan" "Error: only macOS and Linux are supported (detected ${os_name}); architecture must be x86-64 or aarch64/arm64"
4646
exit 1

backends/arm/test/ops/test_matmul.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -455,7 +455,7 @@ def test_matmul_u85_INT(test_case: test_case_t):
455455
pipeline.run()
456456

457457

458-
@common.parametrize("test_case", test_suite | test_suite_fp16)
458+
@common.parametrize("test_case", test_suite | test_suite_fp16 | test_suite_bf16)
459459
@common.SkipIfNoModelConverter
460460
def test_matmul_vgf_no_quant(test_case: test_case_t):
461461
test_data = test_case()

backends/arm/test/runner_utils.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -243,8 +243,7 @@ def is_concrete_shape(shape_like) -> bool:
243243
return all(isinstance(dim, numbers.Integral) for dim in shape_like)
244244

245245
def to_torch_tensor() -> torch.Tensor:
246-
if array.dtype.type is np.void:
247-
# If dtype is void, "cheat" and use the output_tensor dtype.
246+
if output_tensor.dtype == torch.bfloat16 or array.dtype.type is np.void:
248247
return torch.frombuffer(array, dtype=output_tensor.dtype)
249248
return torch.from_numpy(array)
250249

backends/arm/test/tester/test_pipeline.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1258,13 +1258,15 @@ def __init__(
12581258
):
12591259
if tosa_spec is None:
12601260
if tosa_version is None:
1261-
tosa_spec = VgfCompileSpec().tosa_spec
1262-
else:
1263-
if tosa_extensions is None:
1261+
tosa_version = str(VgfCompileSpec().tosa_spec)
1262+
if tosa_extensions is None:
1263+
if "FP" in tosa_version:
1264+
tosa_extensions = ["bf16"]
1265+
else:
12641266
tosa_extensions = []
1265-
tosa_spec = TosaSpecification.create_from_string(
1266-
tosa_version + "".join([f"+{ext}" for ext in tosa_extensions])
1267-
)
1267+
tosa_spec = TosaSpecification.create_from_string(
1268+
tosa_version + "".join([f"+{ext}" for ext in tosa_extensions])
1269+
)
12681270
elif isinstance(tosa_spec, str):
12691271
tosa_spec = TosaSpecification.create_from_string(tosa_spec)
12701272
compile_spec = common.get_vgf_compile_spec(
Submodule Vulkan-Headers updated 69 files

0 commit comments

Comments
 (0)