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

Method intersperse

climart/data_transform/transforms.py:394–413  ·  view source on GitHub ↗
(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)
                    )

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected