MCPcopy Create free account
hub / github.com/NVIDIA/DreamDojo / ComposedModalityTransform

Class ComposedModalityTransform

groot_dreams/data/transform/base.py:78–119  ·  view source on GitHub ↗

Compose multiple modality transforms.

Source from the content-addressed store, hash-verified

76
77
78class ComposedModalityTransform(ModalityTransform):
79 """Compose multiple modality transforms."""
80
81 transforms: list[ModalityTransform] = Field(..., description="The transforms to compose.")
82 apply_to: list[str] = Field(
83 default_factory=list, description="Will be ignored for composed transforms."
84 )
85 training: bool = Field(
86 default=True, description="Whether to apply the transform in training mode."
87 )
88
89 model_config = ConfigDict(arbitrary_types_allowed=True, from_attributes=True)
90
91 def set_metadata(self, dataset_metadata: DatasetMetadata):
92 for transform in self.transforms:
93 transform.set_metadata(dataset_metadata)
94
95 def apply(self, data: dict[str, Any]) -> dict[str, Any]:
96 for i, transform in enumerate(self.transforms):
97 try:
98 data = transform(data)
99 except Exception as e:
100 raise ValueError(f"Error applying transform {i} to data: {e}") from e
101 return data
102
103 def unapply(self, data: dict[str, Any]) -> dict[str, Any]:
104 for i, transform in enumerate(reversed(self.transforms)):
105 if isinstance(transform, InvertibleModalityTransform):
106 try:
107 data = transform.unapply(data)
108 except Exception as e:
109 step = len(self.transforms) - i - 1
110 raise ValueError(f"Error unapplying transform {step} to data: {e}") from e
111 return data
112
113 def train(self):
114 for transform in self.transforms:
115 transform.train()
116
117 def eval(self):
118 for transform in self.transforms:
119 transform.eval()

Callers 2

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected