ggml-webgpu: add support for f16 repeat (#26307)
This commit is contained in:
@@ -2774,6 +2774,10 @@ class ggml_webgpu_shader_lib {
|
|||||||
defines.push_back("TYPE_F32");
|
defines.push_back("TYPE_F32");
|
||||||
variant += "_f32";
|
variant += "_f32";
|
||||||
break;
|
break;
|
||||||
|
case GGML_TYPE_F16:
|
||||||
|
defines.push_back("TYPE_F16");
|
||||||
|
variant += "_f16";
|
||||||
|
break;
|
||||||
case GGML_TYPE_I32:
|
case GGML_TYPE_I32:
|
||||||
defines.push_back("TYPE_I32");
|
defines.push_back("TYPE_I32");
|
||||||
variant += "_i32";
|
variant += "_i32";
|
||||||
|
|||||||
@@ -4290,7 +4290,8 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
|||||||
supports_op = (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_I32);
|
supports_op = (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_I32);
|
||||||
break;
|
break;
|
||||||
case GGML_OP_REPEAT:
|
case GGML_OP_REPEAT:
|
||||||
supports_op = (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_I32 || src0->type == GGML_TYPE_I16);
|
supports_op = (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_I32 ||
|
||||||
|
src0->type == GGML_TYPE_I16);
|
||||||
break;
|
break;
|
||||||
case GGML_OP_CPY:
|
case GGML_OP_CPY:
|
||||||
case GGML_OP_CONT:
|
case GGML_OP_CONT:
|
||||||
|
|||||||
@@ -27,6 +27,9 @@ struct Params {
|
|||||||
#ifdef TYPE_I32
|
#ifdef TYPE_I32
|
||||||
#define DataType i32
|
#define DataType i32
|
||||||
#endif
|
#endif
|
||||||
|
#ifdef TYPE_F16
|
||||||
|
#define DataType f16
|
||||||
|
#endif
|
||||||
#ifdef TYPE_I16
|
#ifdef TYPE_I16
|
||||||
// same size (16-bit) is sufficient for repeat
|
// same size (16-bit) is sufficient for repeat
|
||||||
#define DataType f16
|
#define DataType f16
|
||||||
|
|||||||
@@ -8511,6 +8511,7 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
|
|||||||
test_cases.emplace_back(new test_repeat(GGML_TYPE_F32, {10, 5, 4, ne3}, {1, 2, 1, 1}));
|
test_cases.emplace_back(new test_repeat(GGML_TYPE_F32, {10, 5, 4, ne3}, {1, 2, 1, 1}));
|
||||||
test_cases.emplace_back(new test_repeat(GGML_TYPE_F32, {10, 5, 4, ne3}, {1, 1, 2, 1}));
|
test_cases.emplace_back(new test_repeat(GGML_TYPE_F32, {10, 5, 4, ne3}, {1, 1, 2, 1}));
|
||||||
test_cases.emplace_back(new test_repeat(GGML_TYPE_F32, {10, 5, 4, ne3}, {1, 1, 1, 2}));
|
test_cases.emplace_back(new test_repeat(GGML_TYPE_F32, {10, 5, 4, ne3}, {1, 1, 1, 2}));
|
||||||
|
test_cases.emplace_back(new test_repeat(GGML_TYPE_F16, {10, 5, 4, ne3}, {2, 1, 1, 1}));
|
||||||
test_cases.emplace_back(new test_repeat(GGML_TYPE_I32, {10, 5, 4, ne3}, {2, 1, 1, 1}));
|
test_cases.emplace_back(new test_repeat(GGML_TYPE_I32, {10, 5, 4, ne3}, {2, 1, 1, 1}));
|
||||||
test_cases.emplace_back(new test_repeat(GGML_TYPE_I16, {10, 5, 4, ne3}, {1, 1, 1, 2}));
|
test_cases.emplace_back(new test_repeat(GGML_TYPE_I16, {10, 5, 4, ne3}, {1, 1, 1, 2}));
|
||||||
test_cases.emplace_back(new test_repeat(GGML_TYPE_BF16, {10, 5, 4, ne3}, {2, 1, 1, 1}));
|
test_cases.emplace_back(new test_repeat(GGML_TYPE_BF16, {10, 5, 4, ne3}, {2, 1, 1, 1}));
|
||||||
|
|||||||
Reference in New Issue
Block a user