For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/get-started.md).
Mojo function
selective_scan_fwd_shape
def selective_scan_fwd_shape[dtype: DType](u: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=u.static_spec], delta: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=delta.static_spec], A: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=A.static_spec], B: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=B.static_spec], C: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=C.static_spec], D: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=D.static_spec], z: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=z.static_spec], delta_bias: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=delta_bias.static_spec]) -> IndexList[Int(3)]
Returns the output shape for the selective_scan_fwd op.
The output of a selective scan forward pass has the same shape as the
input sequence u: (batch, dim, seqlen).
Parameters:
- βdtype (
DType): Element type of the input tensors (inferred).
Args:
- βu (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=u.static_spec]): Input sequence tensor with shape(batch, dim, seqlen). - βdelta (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=delta.static_spec]): Time-step tensor with shape(batch, dim, seqlen). - βA (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=A.static_spec]): State transition matrix with shape(dim, dstate). - βB (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=B.static_spec]): Input projection with shape(batch, n_groups, dstate, seqlen). - βC (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=C.static_spec]): Output projection with shape(batch, n_groups, dstate, seqlen). - βD (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=D.static_spec]): Skip connection with shape(dim,). - βz (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=z.static_spec]): Gating tensor with shape(batch, dim, seqlen). - βdelta_bias (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=delta_bias.static_spec]): Delta bias with shape(dim,).
Returns:
IndexList[Int(3)]: The shape of the output tensor (batch, dim, seqlen).
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!