slice_tensor
(tensor, start, end)
| 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 |
no test coverage detected