| 103 | } |
| 104 | |
| 105 | std::vector<matrix_mul::TestArg> matrix_mul::get_matmul_args_split_k() { |
| 106 | std::vector<TestArg> args = get_matmul_args(); |
| 107 | for (auto iter = args.begin(); iter < args.end();) { |
| 108 | if (iter->k <= iter->n) { |
| 109 | iter = args.erase(iter); |
| 110 | } else { |
| 111 | iter++; |
| 112 | } |
| 113 | } |
| 114 | return args; |
| 115 | } |
| 116 | |
| 117 | std::vector<matrix_mul::TestArg> matrix_mul::get_batched_matmul_args_mask( |
| 118 | uint8_t mask) { |