(g, shapes, dtype)
| 298 | |
| 299 | |
| 300 | def get_random_torch_batch(g, shapes, dtype): |
| 301 | is_fp = torch.is_floating_point(torch.tensor([], dtype=dtype)) |
| 302 | if is_fp: |
| 303 | return [torch.rand((shape), generator=g, dtype=dtype) for shape in shapes] |
| 304 | else: |
| 305 | iinfo = torch.iinfo(dtype) |
| 306 | dtype_min, dtype_max = iinfo.min, iinfo.max |
| 307 | return [ |
| 308 | torch.randint(dtype_min, dtype_max, shape, generator=g, dtype=dtype) for shape in shapes |
| 309 | ] |
| 310 | |
| 311 | |
| 312 | def get_sliced_torch_case(case_name): |
no outgoing calls
no test coverage detected