MCPcopy Create free account
hub / github.com/PaddlePaddle/FastDeploy / fn

Function fn

fastdeploy/model_executor/models/tp_utils.py:215–325  ·  view source on GitHub ↗

func

(x, is_column=True)

Source from the content-addressed store, hash-verified

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

Callers 1

wrapped_fnFunction · 0.50

Calls 6

slice_tensorFunction · 0.85
split_tensorFunction · 0.85
split_or_merge_qkv_funcFunction · 0.85
funcFunction · 0.85
popMethod · 0.80
get_default_dtypeMethod · 0.80

Tested by

no test coverage detected