func
(x, is_column=True)
| 213 | """ |
| 214 | |
| 215 | def fn(x, is_column=True): |
| 216 | """func""" |
| 217 | |
| 218 | def get_shape(tensor): |
| 219 | """get_shape""" |
| 220 | return tensor.get_shape() if hasattr(tensor, "get_shape") else tensor.shape |
| 221 | |
| 222 | def slice_tensor(tensor, start, end): |
| 223 | """slice_tensor""" |
| 224 | shape = get_shape(tensor) |
| 225 | if len(shape) == 1: |
| 226 | return tensor[start:end] |
| 227 | elif is_column: |
| 228 | return tensor[..., start:end] |
| 229 | else: |
| 230 | return tensor[start:end, ...] |
| 231 | |
| 232 | q_end = num_attention_heads * head_dim |
| 233 | k_end = q_end + num_key_value_heads * head_dim |
| 234 | v_end = k_end + num_key_value_heads * head_dim |
| 235 | |
| 236 | q = slice_tensor(x, 0, q_end) |
| 237 | k = slice_tensor(x, q_end, k_end) |
| 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 = ( |
| 263 | num_key_value_heads < tensor_model_parallel_size and tensor_model_parallel_size % num_key_value_heads == 0 |
| 264 | ) |
| 265 | repeat_num = tensor_model_parallel_size // num_key_value_heads if repeat_kv else 1 |
| 266 | if repeat_kv: |
| 267 | k_list = split_tensor(k, num_key_value_heads) |
| 268 | v_list = split_tensor(v, num_key_value_heads) |
| 269 | else: |
| 270 | k_list = split_tensor(k, tensor_model_parallel_size) |
| 271 | v_list = split_tensor(v, tensor_model_parallel_size) |
| 272 |
no test coverage detected