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

Function split_tensor

fastdeploy/model_executor/models/tp_utils.py:240–259  ·  view source on GitHub ↗

split_tensor

(tensor, degree)

Source from the content-addressed store, hash-verified

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 = (

Callers 2

gqa_qkv_split_funcMethod · 0.85
fnFunction · 0.85

Calls 3

get_shapeFunction · 0.85
slice_tensorFunction · 0.85
splitMethod · 0.80

Tested by

no test coverage detected