MCPcopy Create free account
hub / github.com/RolnickLab/climart / ColumnPreprocesser

Class ColumnPreprocesser

climart/data_transform/transforms.py:357–433  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

355
356
357class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected