| 1135 | self.assertTrue(mx.array_equal(grad_ind, mx.zeros(ind.shape))) |
| 1136 | |
| 1137 | def test_setitem(self): |
| 1138 | a = mx.array(0) |
| 1139 | a[None] = 1 |
| 1140 | self.assertEqual(a.item(), 1) |
| 1141 | |
| 1142 | a = mx.array([1, 2, 3]) |
| 1143 | a[0] = 2 |
| 1144 | self.assertEqual(a.tolist(), [2, 2, 3]) |
| 1145 | |
| 1146 | a[-1] = 2 |
| 1147 | self.assertEqual(a.tolist(), [2, 2, 2]) |
| 1148 | |
| 1149 | a[np.int64(1)] = 9 |
| 1150 | self.assertEqual(a.tolist(), [2, 9, 2]) |
| 1151 | |
| 1152 | a[0] = mx.array([[[1]]]) |
| 1153 | self.assertEqual(a.tolist(), [1, 9, 2]) |
| 1154 | |
| 1155 | a[:] = 0 |
| 1156 | self.assertEqual(a.tolist(), [0, 0, 0]) |
| 1157 | |
| 1158 | a[None] = 1 |
| 1159 | self.assertEqual(a.tolist(), [1, 1, 1]) |
| 1160 | |
| 1161 | a[0:1] = 2 |
| 1162 | self.assertEqual(a.tolist(), [2, 1, 1]) |
| 1163 | |
| 1164 | a[0:2] = 3 |
| 1165 | self.assertEqual(a.tolist(), [3, 3, 1]) |
| 1166 | |
| 1167 | a[0:3] = 4 |
| 1168 | self.assertEqual(a.tolist(), [4, 4, 4]) |
| 1169 | |
| 1170 | a[0:1] = mx.array(0) |
| 1171 | self.assertEqual(a.tolist(), [0, 4, 4]) |
| 1172 | |
| 1173 | a[0:1] = mx.array([1]) |
| 1174 | self.assertEqual(a.tolist(), [1, 4, 4]) |
| 1175 | |
| 1176 | with self.assertRaises(ValueError): |
| 1177 | a[0:1] = mx.array([2, 3]) |
| 1178 | |
| 1179 | a[0:2] = mx.array([2, 2]) |
| 1180 | self.assertEqual(a.tolist(), [2, 2, 4]) |
| 1181 | |
| 1182 | a[:] = mx.array([[[[1, 1, 1]]]]) |
| 1183 | self.assertEqual(a.tolist(), [1, 1, 1]) |
| 1184 | |
| 1185 | # Array slices |
| 1186 | def check_slices(arr_np, update_np, *idx_np): |
| 1187 | arr_mlx = mx.array(arr_np) |
| 1188 | update_mlx = mx.array(update_np) |
| 1189 | idx_mlx = [ |
| 1190 | mx.array(idx) if isinstance(idx, np.ndarray) else idx for idx in idx_np |
| 1191 | ] |
| 1192 | if len(idx_np) > 1: |
| 1193 | idx_np = tuple(idx_np) |
| 1194 | idx_mlx = tuple(idx_mlx) |