MCPcopy Create free account
hub / github.com/ml-explore/mlx / test_setitem

Method test_setitem

python/tests/test_array.py:1137–1392  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 3

itemMethod · 0.80
arrayMethod · 0.60
sliceFunction · 0.50

Tested by

no test coverage detected