(num_iters, operator_fn, fixed_params, input_params, input_data: ArgData, rng)
| 265 | |
| 266 | |
| 267 | def _test_seq_input(num_iters, operator_fn, fixed_params, input_params, input_data: ArgData, rng): |
| 268 | @pipeline_def |
| 269 | def pipeline(args_data: List[ArgData]): |
| 270 | pos_args = [arg_data for arg_data in args_data if arg_data.desc.is_positional_arg] |
| 271 | pos_nodes = [None] * len(pos_args) |
| 272 | for arg_data in pos_args: |
| 273 | assert 0 <= arg_data.desc.name < len(pos_nodes) |
| 274 | assert pos_nodes[arg_data.desc.name] is None |
| 275 | pos_nodes[arg_data.desc.name] = arg_data_node(arg_data) |
| 276 | named_args = [arg_data for arg_data in args_data if not arg_data.desc.is_positional_arg] |
| 277 | arg_nodes = {arg_data.desc.name: arg_data_node(arg_data) for arg_data in named_args} |
| 278 | output = operator_fn(*pos_nodes, **fixed_params, **arg_nodes) |
| 279 | return output |
| 280 | |
| 281 | assert num_iters >= len(input_data.data) |
| 282 | max_batch_size = max(len(batch) for batch in input_data.data) |
| 283 | |
| 284 | params_provider = ( |
| 285 | input_params |
| 286 | if isinstance(input_params, ParamsProviderBase) |
| 287 | else ParamsProvider(input_params) |
| 288 | ) |
| 289 | params_provider.setup(input_data, fixed_params, rng) |
| 290 | args_data = params_provider.compute_params() |
| 291 | seq_pipe = pipeline( |
| 292 | args_data=[input_data, *args_data], batch_size=max_batch_size, num_threads=4, device_id=0 |
| 293 | ) |
| 294 | unfolded_input = params_provider.unfold_input() |
| 295 | expanded_args_data = params_provider.expand_params() |
| 296 | max_uf_batch_size = max(len(batch) for batch in unfolded_input.data) |
| 297 | baseline_pipe = pipeline( |
| 298 | args_data=[unfolded_input, *expanded_args_data], |
| 299 | batch_size=max_uf_batch_size, |
| 300 | num_threads=4, |
| 301 | device_id=0, |
| 302 | ) |
| 303 | |
| 304 | for _ in range(num_iters): |
| 305 | (seq_batch,) = seq_pipe.run() |
| 306 | (baseline_batch,) = baseline_pipe.run() |
| 307 | unfolded_layout = params_provider.unfold_output_layout(seq_batch.layout()) |
| 308 | assert ( |
| 309 | unfolded_layout == baseline_batch.layout() |
| 310 | ), f"'{unfolded_layout}' == '{baseline_batch.layout()}'" |
| 311 | batch = params_provider.unfold_output(as_batch(seq_batch)) |
| 312 | baseline_batch = as_batch(baseline_batch) |
| 313 | assert len(batch) == len(baseline_batch) |
| 314 | check_batch(batch, baseline_batch, len(batch)) |
| 315 | |
| 316 | |
| 317 | def get_input_arg_per_sample(input_data, param_cb, rng): |
nothing calls this directly
no test coverage detected