(rng)
| 45 | op_instance = None |
| 46 | |
| 47 | def generate(rng): |
| 48 | nonlocal op_instance |
| 49 | # Create operator or use functional API |
| 50 | if api_type == "ops": |
| 51 | if op_instance is None: |
| 52 | op_instance = ops[opname](device=device_type, max_batch_size=batch_size) |
| 53 | result1 = op_instance(batch_size=batch_size, rng=rng, **op_args[opname]) |
| 54 | else: |
| 55 | result1 = fn[opname]( |
| 56 | batch_size=batch_size, rng=rng, device=device_type, **op_args[opname] |
| 57 | ) |
| 58 | |
| 59 | # Verify result type and shape |
| 60 | if batch_size is not None: |
| 61 | assert isinstance(result1, ndd.Batch), f"Expected Batch, got {type(result1)}" |
| 62 | result1_np = asnumpy(result1) |
| 63 | assert result1_np.shape == ( |
| 64 | batch_size, |
| 65 | 10, |
| 66 | ), f"Expected shape ({batch_size}, 10), got {result1_np.shape}" |
| 67 | else: |
| 68 | assert isinstance(result1, ndd.Tensor), f"Expected Tensor, got {type(result1)}" |
| 69 | result1_np = asnumpy(result1) |
| 70 | assert result1_np.shape == (10,), f"Expected shape (10,), got {result1_np.shape}" |
| 71 | return result1_np |
| 72 | |
| 73 | rng1 = ndd.random.RNG(seed=1234) |
| 74 | rng2 = ndd.random.RNG(seed=1234) |
no test coverage detected