(self,
name_to_array: Optional[Dict[str, Tensor]] = None,
global_node: Optional[Tensor] = None, # shape (b, #feats)
levels: Optional[Tensor] = None, # shape (b, #levels, #feats)
layers: Optional[Tensor] = None # shape (b, #layers, #feats)
)
| 392 | return preprocesser |
| 393 | |
| 394 | def intersperse(self, |
| 395 | name_to_array: Optional[Dict[str, Tensor]] = None, |
| 396 | global_node: Optional[Tensor] = None, # shape (b, #feats) |
| 397 | levels: Optional[Tensor] = None, # shape (b, #levels, #feats) |
| 398 | layers: Optional[Tensor] = None # shape (b, #layers, #feats) |
| 399 | ) -> Tensor: |
| 400 | global_node, levels, layers = self.get_data_types(name_to_array, global_node, levels, layers) |
| 401 | |
| 402 | if global_node.shape[-1] != layers.shape[-1] or levels.shape[-1] != layers.shape[-1]: |
| 403 | raise ValueError("Expected all node types to have same dimensions. Project them first or pad them instead!") |
| 404 | |
| 405 | batch_size, _, n_feats = levels.shape |
| 406 | |
| 407 | interspersed_data = torch.empty((batch_size, self.n_nodes, n_feats)) |
| 408 | interspersed_data[:, self.GLOBAL_NODE, :] = global_node |
| 409 | interspersed_data[:, self.LEVEL_NODES, :] = levels |
| 410 | interspersed_data[:, self.LAYER_NODES, :] = layers |
| 411 | |
| 412 | interspersed_data = interspersed_data.to(global_node.device) |
| 413 | return interspersed_data |
| 414 | |
| 415 | def intersperse_no_levels(self, |
| 416 | name_to_array: Optional[Dict[str, Tensor]] = None, |
nothing calls this directly
no outgoing calls
no test coverage detected