(self, shape: Tuple[int, ...], tensor_type: str = "np", **kwargs)
| 7 | |
| 8 | class NoiseProcess: |
| 9 | def __init__(self, shape: Tuple[int, ...], tensor_type: str = "np", **kwargs): |
| 10 | self.shape = shape |
| 11 | self.x = None |
| 12 | if tensor_type == "np": |
| 13 | self.randn = np.random.randn |
| 14 | self.zeros = np.zeros |
| 15 | elif tensor_type.startswith("torch_"): |
| 16 | # syntax: torch_<device>, torch_cpu, torch_cuda |
| 17 | self.randn = lambda *shape: torch.randn(*shape, device=tensor_type.split("_")[-1]) |
| 18 | self.zeros = lambda *shape: torch.zeros(*shape, device=tensor_type.split("_")[-1]) |
| 19 | else: |
| 20 | raise ValueError(f"Invalid tensor type: {tensor_type}") |
| 21 | |
| 22 | def reset(self): |
| 23 | raise NotImplementedError |
nothing calls this directly
no outgoing calls
no test coverage detected