@staticmethod
def forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:
"""Override in subclass to define the forward computation on local tensors."""
raise NotImplementedError("Subclasses must implement forward()")
@staticmethod
def backward(ctx: Any, *grad_outputs: Any) -> Any:
"""Override in subclass to define the backward computation on local tensors."""
raise NotImplementedError("Subclasses must implement backward()")
@classmethod
def apply(cls, *args: Any, **kwargs: Any) -> Any:
"""Execute the function, routing to distributed dispatch when DTensor inputs are present.