(
self,
inputs: Tuple[Any, ...],
kwargs: Optional[Dict[str, Any]],
device_ids: Sequence[Union[int, torch.device]],
)
| 189 | return replicate(module, device_ids, not torch.is_grad_enabled()) |
| 190 | |
| 191 | def scatter( |
| 192 | self, |
| 193 | inputs: Tuple[Any, ...], |
| 194 | kwargs: Optional[Dict[str, Any]], |
| 195 | device_ids: Sequence[Union[int, torch.device]], |
| 196 | ) -> Any: |
| 197 | return scatter_kwargs(inputs, kwargs, device_ids, dim=self.dim) |
| 198 | |
| 199 | def parallel_apply(self, replicas: Sequence[T], inputs: Sequence[Any], kwargs: Any) -> List[Any]: |
| 200 | return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) |
no test coverage detected