| 355 | |
| 356 | |
| 357 | class ColumnPreprocesser: |
| 358 | ONLY_LAYER_NODES = ['duplication', 'graph_net_layer_nodes', 'identity'] |
| 359 | ONLY_LEVEL_NODES = ['graph_net_level_nodes'] |
| 360 | |
| 361 | def __init__(self, |
| 362 | preprocessing: str, |
| 363 | exp_type: str, |
| 364 | projector_hidden_dim: int = 128, # only if preprocessing == 'mlp' |
| 365 | projector_n_layers: int = 1, # only if preprocessing == 'mlp' |
| 366 | projector_net_normalization: str = 'layer_norm', # only if preprocessing == 'mlp' |
| 367 | use_level_features: bool = True, |
| 368 | ): |
| 369 | self.preprocessing_type = preprocessing.lower() |
| 370 | self.projector_hidden_dim = projector_hidden_dim |
| 371 | self.projector_n_layers = projector_n_layers |
| 372 | self.projector_net_normalization = projector_net_normalization |
| 373 | self.use_level_features = use_level_features |
| 374 | |
| 375 | def get_preprocesser(self, batched: bool = False, verbose: bool = True): |
| 376 | if self.preprocessing_type in ['mlp', 'mlp_projection']: |
| 377 | in_dims = self.input_dims.copy() |
| 378 | if not self.use_level_features: |
| 379 | in_dims.pop(LEVELS) |
| 380 | preprocesser = FeatureProjector( |
| 381 | input_name_to_feature_dim=in_dims, |
| 382 | projector_n_layers=self.projector_n_layers, |
| 383 | projection_dim=self.projector_hidden_dim, |
| 384 | projector_activation_func='Gelu', |
| 385 | projector_net_normalization=self.projector_net_normalization, |
| 386 | output_normalization=True, |
| 387 | output_activation_function=False, |
| 388 | projections_aggregation=self.intersperse if self.use_level_features else self.intersperse_no_levels) |
| 389 | self.out_dim = self.projector_hidden_dim |
| 390 | self.as_string = f'All var types are MLP-projected to a {self.projector_hidden_dim} hidden dimension' |
| 391 | |
| 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 |
nothing calls this directly
no outgoing calls
no test coverage detected