Team Ai
Modelpublic

nvidia/C-RADIOv4-H

sourceHugging Faceotherupdated 8mo agoView on Hugging Face
85likes27kdownloads
adaptor_base.py51 linesDownload Raw Back to root
1# Copyright (c) 2024, NVIDIA CORPORATION.  All rights reserved.2#3# NVIDIA CORPORATION and its licensors retain all intellectual property4# and proprietary rights in and to this software, related documentation5# and any modifications thereto.  Any use, reproduction, disclosure or6# distribution of this software and related documentation without an express7# license agreement from NVIDIA CORPORATION is strictly prohibited.8from argparse import Namespace9from typing import NamedTuple, Optional10 11import torch12from torch import nn13import torch.nn.functional as F14 15 16class AdaptorInput(NamedTuple):17    images: torch.Tensor18    summary: torch.Tensor19    features: torch.Tensor20    feature_fmt: str21    patch_size: int22 23 24class RadioOutput(NamedTuple):25    summary: torch.Tensor26    features: torch.Tensor27 28    def to(self, *args, **kwargs):29        return RadioOutput(30            self.summary.to(*args, **kwargs) if self.summary is not None else None,31            self.features.to(*args, **kwargs) if self.features is not None else None,32        )33 34 35class AdaptorModuleBase(nn.Module):36    def __init__(37        self,38        requires_summary_and_spatial: bool,39        handles_summary_and_spatial: bool = False40    ) -> None:41        super().__init__()42        self.requires_summary_and_spatial = requires_summary_and_spatial43        self.handles_summary_and_spatial = handles_summary_and_spatial44 45        assert not handles_summary_and_spatial or requires_summary_and_spatial, "If handles summary and spatial, must require it too!"46 47 48class AdaptorBase(nn.Module):49    def forward(self, input: AdaptorInput) -> RadioOutput:50        raise NotImplementedError("Subclasses must implement this!")51