MCPcopy Create free account
hub / github.com/apple/axlearn / _get_visual_features

Method _get_visual_features

axlearn/vision/virtex.py:116–155  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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
158class VirTexModel(ImageBackboneModelMixin, BaseLayer):

Callers 1

embed_imageMethod · 0.80

Calls 3

get_recursivelyFunction · 0.90
getMethod · 0.80
get_module_outputsMethod · 0.45

Tested by

no test coverage detected