(self)
| 884 | |
| 885 | @unittest.skipIf(not mx.is_available(mx.gpu), "No GPU available") |
| 886 | def test_custom_kernel_helper(self): |
| 887 | if mx.metal.is_available(): |
| 888 | header = """ |
| 889 | template <typename T> |
| 890 | T do_exp(T x) { |
| 891 | return metal::precise::exp(x); |
| 892 | } |
| 893 | """ |
| 894 | source = """ |
| 895 | uint elem = thread_position_in_grid.x; |
| 896 | out1[elem] = do_exp(a[elem]); |
| 897 | """ |
| 898 | custom_kernel = mx.fast.metal_kernel |
| 899 | elif mx.cuda.is_available(): |
| 900 | header = """ |
| 901 | template <typename T> |
| 902 | __device__ T do_exp(T x) { |
| 903 | return exp(x); |
| 904 | } |
| 905 | """ |
| 906 | source = """ |
| 907 | auto elem = cooperative_groups::this_grid().thread_rank(); |
| 908 | out1[elem] = do_exp(a[elem]); |
| 909 | """ |
| 910 | custom_kernel = mx.fast.cuda_kernel |
| 911 | |
| 912 | mx.random.seed(7) |
| 913 | a = mx.random.normal(shape=(2, 2)) |
| 914 | kernel = custom_kernel( |
| 915 | name="helper", |
| 916 | input_names=["a"], |
| 917 | output_names=["out1"], |
| 918 | header=header, |
| 919 | source=source, |
| 920 | ) |
| 921 | out = kernel( |
| 922 | inputs=[a], |
| 923 | grid=(4, 1, 1), |
| 924 | threadgroup=(2, 1, 1), |
| 925 | output_shapes=[(2, 2)], |
| 926 | output_dtypes=[mx.float32], |
| 927 | stream=mx.gpu, |
| 928 | ) |
| 929 | self.assertTrue(mx.allclose(out[0], mx.exp(a))) |
| 930 | |
| 931 | @unittest.skipIf(not mx.is_available(mx.gpu), "No GPU available") |
| 932 | def test_custom_kernel_attributes(self): |
nothing calls this directly
no test coverage detected