MCPcopy Create free account
hub / github.com/PolymathicAI/AstroCLIP / forward

Method forward

astroclip/models/astroclip.py:134–154  ·  view source on GitHub ↗
(
        self,
        image_features: torch.FloatTensor,
        spectrum_features: torch.FloatTensor,
        logit_scale: float,
        output_dict: bool = False,
    )

Source from the content-addressed store, hash-verified

132 return logits_per_image, logits_per_image.T
133
134 def forward(
135 self,
136 image_features: torch.FloatTensor,
137 spectrum_features: torch.FloatTensor,
138 logit_scale: float,
139 output_dict: bool = False,
140 ) -> torch.FloatTensor:
141 # Get the logits for the image and spectrum features
142 logits_per_image, logits_per_spectrum = self.get_logits(
143 image_features, spectrum_features, logit_scale
144 )
145
146 # Calculate the contrastive loss
147 labels = torch.arange(
148 logits_per_image.shape[0], device=image_features.device, dtype=torch.long
149 )
150 total_loss = (
151 F.cross_entropy(logits_per_image, labels)
152 + F.cross_entropy(logits_per_spectrum, labels)
153 ) / 2
154 return {"contrastive_loss": total_loss} if output_dict else total_loss
155
156
157class ImageHead(nn.Module):

Callers

nothing calls this directly

Calls 1

get_logitsMethod · 0.95

Tested by

no test coverage detected