Compose multiple modality transforms.
| 76 | |
| 77 | |
| 78 | class 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() |
no outgoing calls
no test coverage detected