(size: dict, dtype: torch.dtype, device: str, seed: int = 42)
| 216 | |
| 217 | |
| 218 | def gen_layernorm_inputs(size: dict, dtype: torch.dtype, device: str, seed: int = 42) -> dict: |
| 219 | torch.manual_seed(seed) |
| 220 | batch, dim = size["batch"], size["dim"] |
| 221 | x = torch.randn(batch, dim, device=device, dtype=dtype) |
| 222 | weight = torch.ones(dim, device=device, dtype=dtype) |
| 223 | bias = torch.zeros(dim, device=device, dtype=dtype) |
| 224 | return {"x": x, "weight": weight, "bias": bias} |
| 225 | |
| 226 | |
| 227 | def gen_flash_attention_inputs(size: dict, dtype: torch.dtype, device: str, seed: int = 42) -> dict: |
nothing calls this directly
no outgoing calls
no test coverage detected