transforms `BatchGraph`. Default transformation is encoding nodes and edges(if exist) feature to embedding. Args: transform_func: A function that takes in an `BatchGraph` object and returns a transformed version.
(self, transform_func=None)
| 116 | setattr(self, key, value) |
| 117 | |
| 118 | def transform(self, transform_func=None): |
| 119 | """transforms `BatchGraph`. Default transformation is encoding |
| 120 | nodes and edges(if exist) feature to embedding. |
| 121 | Args: |
| 122 | transform_func: A function that takes in an `BatchGraph` object |
| 123 | and returns a transformed version. |
| 124 | """ |
| 125 | def transform_feat(feat, schema): |
| 126 | feat_handler = FeatureHandler(schema[0], schema[1]) |
| 127 | return feat_handler.forward(feat) |
| 128 | |
| 129 | if self.node_schema is None: |
| 130 | return self |
| 131 | |
| 132 | node = Data(self.nodes.ids, |
| 133 | self.nodes.int_attrs, |
| 134 | self.nodes.float_attrs, |
| 135 | self.nodes.string_attrs) |
| 136 | node_tensor = transform_feat(node, |
| 137 | [self.node_schema[0], self.node_schema[1].feature_spec]) |
| 138 | |
| 139 | edge_tensor = None |
| 140 | if self.edge_schema is not None: |
| 141 | edge = Data(self.edges.ids, |
| 142 | self.edges.int_attrs, |
| 143 | self.edges.float_attrs, |
| 144 | self.edges.string_attrs) |
| 145 | edge_tensor = transform_feat(edge, |
| 146 | [self.edge_schema[0], self.edge_schema[1].feature_spec]) |
| 147 | |
| 148 | graph = BatchGraph(self.edge_index, |
| 149 | node_tensor, self.node_schema, self.graph_node_offsets, |
| 150 | edge_tensor, self.edge_schema, self.graph_edge_offsets, |
| 151 | additional_keys=self.additional_keys) |
| 152 | for key in self.additional_keys: |
| 153 | graph[key] = self[key] |
| 154 | return graph |
| 155 | |
| 156 | def to_graphs(self): |
| 157 | """reconstructs `SubGraph`s.""" |
nothing calls this directly
no test coverage detected