qr0, qr1 = split_half(kwargs.get("query_rope"), seq_dim)
key_rope = kwargs.get("key_rope")
keys = unfold_prefix_pair(key, seq_shard_id, seq_shards, seq_dim)
values = unfold_prefix_pair(value, seq_shard_id, seq_shards, seq_dim)
key_ropes = unfold_prefix_pair(key_rope, seq_shard_id, seq_shards, seq_dim)
def _call(block_id, q, topk, q_rope):
kw = dict(kwargs)
if "query_rope" in kw:
kw["query_rope"] = q_rope
if "key_rope" in kw:
kw["key_rope"] = key_ropes[block_id]
if tnd:
q_len, k_len = tnd_block_seq_lens(
kw.get("actual_seq_lengths_query"), kw.get("actual_seq_lengths_kv"),
key.shape[0], seq_shard_id, seq_shards, block_id)
kw["actual_seq_lengths_query"] = q_len
kw["actual_seq_lengths_kv"] = k_len
return func(q, keys[block_id], values[block_id], topk, *rest, **kw)
out0 = _call(0, q0, t0, qr0)
out1 = _call(1, q1, t1, qr1)
if not isinstance(out0, (tuple, list)):
return _cat_pair(out0, out1, seq_dim)
stitched = [_cat_pair(out0[0], out1[0], seq_dim)]
stitched.extend(_cat_pair(a, b, stats_dim) for a, b in zip(out0[1:], out1[1:]))
return _same_sequence_type(out0, stitched)
# ``args`` positions of the MindSpore positional form of the sparse indexer KL loss:
# 0 query, 1 key, 2 query_index, 3 key_index, 4 weights, 5 sparse_indices,