(self, program, feed, scope, data_op_infos)
| 1457 | break |
| 1458 | |
| 1459 | def _pir_feed_data(self, program, feed, scope, data_op_infos): |
| 1460 | # feed var to framework |
| 1461 | feed_target_names = set() |
| 1462 | for data_op_info in data_op_infos: |
| 1463 | feed_target_name = data_op_info[0] |
| 1464 | feed_target_names.add(feed_target_name) |
| 1465 | var_type = data_op_info[1] |
| 1466 | var_shape = data_op_info[2] |
| 1467 | is_persistable = data_op_info[3] |
| 1468 | if feed_target_name not in feed.keys() and is_persistable: |
| 1469 | # If the feed_target_name is not in feed list, but is persistable, maybe it is a optimizer param |
| 1470 | # and don't need feed data. |
| 1471 | continue |
| 1472 | cur_feed = feed[feed_target_name] |
| 1473 | if not isinstance(cur_feed, core.DenseTensor): |
| 1474 | cur_feed = _as_lodtensor(cur_feed, self.place, var_type) |
| 1475 | pir_check_feed_shape_type( |
| 1476 | cur_feed, feed_target_name, var_shape, var_type |
| 1477 | ) |
| 1478 | # the last arg of set_feed_variable has no effect in pir, we pass 0 by default. |
| 1479 | core.set_feed_variable(scope, cur_feed, feed_target_name, 0) |
| 1480 | |
| 1481 | # pop variable which is not found in program |
| 1482 | for feed_name in list(feed.keys()): |
| 1483 | if feed_name not in feed_target_names: |
| 1484 | feed.pop(feed_name) |
| 1485 | warnings.warn( |
| 1486 | f"The value {feed_name} is not found in program. It is not declared or is pruned." |
| 1487 | ) |
| 1488 | |
| 1489 | def _fetch_data(self, fetch_list, fetch_var_name, scope): |
| 1490 | outs = [ |
no test coverage detected