@staticmethod
def backward(ctx, grad0, grad1): # pylint: disable=arguments-differ
"""Merged reverse of both prefixes: the shorter one is a head of the longer one,
so add it into that head and scatter back with a single gather."""
seq_shard_id, seq_shards, full_len, seq_dim, sf, m = ctx.fold_args
long_id = 0 if m[0] >= m[1] else 1
grads = (grad0, grad1)
g_long, g_short = grads[long_id], grads[1 - long_id]
if g_long is None and g_short is None:
return None, None, None, None
if g_long is None:
return fold_prefix_grad(g_short, seq_shard_id, seq_shards, 1 - long_id, full_len, seq_dim), \
None, None, None
parts = []
short_len = m[1 - long_id] * sf
long_len = m[long_id] * sf
if g_short is not None:
parts.append(g_long.narrow(seq_dim, 0, short_len) + g_short)
if long_len > short_len:
parts.append(g_long.narrow(seq_dim, short_len, long_len - short_len))
else:
parts.append(g_long)
zero_shape = list(g_long.shape)
zero_shape[seq_dim] = sf
parts.append(platform.zeros(tuple(zero_shape), dtype=g_long.dtype, device=g_long.device))
g = platform.cat(parts, dim=seq_dim)
g = g.reshape(_chunked_shape(g.shape, seq_dim, m[long_id] + 1))
g = g.index_select(seq_dim, _index_tensor(
build_prefix_scatter_order(seq_shard_id, seq_shards, long_id), g))
out_shape = list(g_long.shape)
out_shape[seq_dim] = full_len
return g.reshape(tuple(out_shape)), None, None, None
def unfold_prefix_pair(x, seq_shard_id: int, seq_shards: int, seq_dim: int = 1):
"""``(unfold_prefix(x, .., 0), unfold_prefix(x, .., 1))`` with one merged backward."""