split_tensor
(tensor, degree)
| 238 | v = slice_tensor(x, k_end, v_end) |
| 239 | |
| 240 | def split_tensor(tensor, degree): |
| 241 | """ |
| 242 | split_tensor |
| 243 | """ |
| 244 | shape = get_shape(tensor) |
| 245 | size = shape[-1] if is_column else shape[0] |
| 246 | block_size = size // degree |
| 247 | if hasattr(tensor, "get_shape"): |
| 248 | return [slice_tensor(tensor, i * block_size, (i + 1) * block_size) for i in range(degree)] |
| 249 | else: |
| 250 | if isinstance(x, paddle.Tensor): |
| 251 | if is_column: |
| 252 | return paddle.split(tensor, degree, axis=-1) |
| 253 | else: |
| 254 | return paddle.split(tensor, degree, axis=0) |
| 255 | else: |
| 256 | if is_column: |
| 257 | return np.split(tensor, degree, axis=-1) |
| 258 | else: |
| 259 | return np.split(tensor, degree, axis=0) |
| 260 | |
| 261 | q_list = split_tensor(q, tensor_model_parallel_size) |
| 262 | repeat_kv = ( |
no test coverage detected