Gets the visual features. The visual features will be extracted from a NestedTensor based on a path specified in `cfg.visual_feature_layer_name`. The NestedTensor can be either `visual_outputs` or self.get_module_outputs()["visual"]. Args: visual_output
(self, visual_outputs: NestedTensor)
| 114 | return self.embed_image(input_batch["image"]) |
| 115 | |
| 116 | def _get_visual_features(self, visual_outputs: NestedTensor) -> Tensor: |
| 117 | """Gets the visual features. |
| 118 | |
| 119 | The visual features will be extracted from a NestedTensor based on a path specified in |
| 120 | `cfg.visual_feature_layer_name`. |
| 121 | |
| 122 | The NestedTensor can be either `visual_outputs` or self.get_module_outputs()["visual"]. |
| 123 | |
| 124 | Args: |
| 125 | visual_outputs: The outputs of the visual layer. |
| 126 | |
| 127 | Returns: |
| 128 | visual_features: The features from the visual backbone. |
| 129 | Shape: (batch, height, width, channels) |
| 130 | |
| 131 | Raises: |
| 132 | ValueError: If visual_feature_layer_name cannot be found. |
| 133 | """ |
| 134 | cfg = self.config |
| 135 | try: |
| 136 | return get_recursively(visual_outputs, cfg.visual_feature_layer_name) |
| 137 | except KeyError: |
| 138 | pass |
| 139 | |
| 140 | try: |
| 141 | return get_recursively( |
| 142 | self.get_module_outputs(), |
| 143 | f"visual/{cfg.visual_feature_layer_name}", |
| 144 | ) |
| 145 | except KeyError: |
| 146 | pass |
| 147 | |
| 148 | def _paths(x): |
| 149 | return jax.tree_util.tree_leaves(tree_paths(x)) |
| 150 | |
| 151 | raise ValueError( |
| 152 | f"Cannot find visual features at {cfg.visual_feature_layer_name}. " |
| 153 | f"visual_outputs={_paths(visual_outputs)}, " |
| 154 | f"module_outputs={_paths(self.get_module_outputs().get('visual'))}" |
| 155 | ) |
| 156 | |
| 157 | |
| 158 | class VirTexModel(ImageBackboneModelMixin, BaseLayer): |
no test coverage detected