Skip to content

clip.marsclip_model

marsclip_model

First tri-modal MarsCLIP model slice with valid-token masking.

ImageViTTower

ImageViTTower(*, image_size: int = 224, patch_size: int = 16, hidden_dim: int = 256, depth: int = 4, num_heads: int = 8, embed_dim: int = 256)

Bases: Module

A lightweight ViT-style image encoder with valid-patch masking.

Source code in src/clip/marsclip_model.py
def __init__(
    self,
    *,
    image_size: int = 224,
    patch_size: int = 16,
    hidden_dim: int = 256,
    depth: int = 4,
    num_heads: int = 8,
    embed_dim: int = 256,
) -> None:
    super().__init__()
    if image_size % patch_size != 0:
        raise ValueError("image_size must be divisible by patch_size.")
    self.image_size = image_size
    self.patch_size = patch_size
    self.num_patches = (image_size // patch_size) ** 2

    self.patch_embed = nn.Conv2d(
        in_channels=3,
        out_channels=hidden_dim,
        kernel_size=patch_size,
        stride=patch_size,
    )
    self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, hidden_dim))
    encoder_layer = nn.TransformerEncoderLayer(
        d_model=hidden_dim,
        nhead=num_heads,
        batch_first=True,
        dim_feedforward=hidden_dim * 4,
        activation="gelu",
    )
    self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth)
    self.proj = nn.Linear(hidden_dim, embed_dim)

TextTransformerTower

TextTransformerTower(*, vocab_size: int, max_length: int = 64, hidden_dim: int = 256, depth: int = 2, num_heads: int = 8, embed_dim: int = 256)

Bases: Module

A lightweight text encoder for raw/expanded rationale text.

Source code in src/clip/marsclip_model.py
def __init__(
    self,
    *,
    vocab_size: int,
    max_length: int = 64,
    hidden_dim: int = 256,
    depth: int = 2,
    num_heads: int = 8,
    embed_dim: int = 256,
) -> None:
    super().__init__()
    self.token_embed = nn.Embedding(vocab_size, hidden_dim)
    self.pos_embed = nn.Parameter(torch.zeros(1, max_length, hidden_dim))
    encoder_layer = nn.TransformerEncoderLayer(
        d_model=hidden_dim,
        nhead=num_heads,
        batch_first=True,
        dim_feedforward=hidden_dim * 4,
        activation="gelu",
    )
    self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth)
    self.proj = nn.Linear(hidden_dim, embed_dim)

GeoContextTower

GeoContextTower(*, geo_dim: int, scale_dim: int, viewing_dim: int, quality_dim: int, hidden_dim: int = 256, embed_dim: int = 256)

Bases: Module

Encode geo, scale, viewing, and quality metadata into one embedding.

Source code in src/clip/marsclip_model.py
def __init__(
    self,
    *,
    geo_dim: int,
    scale_dim: int,
    viewing_dim: int,
    quality_dim: int,
    hidden_dim: int = 256,
    embed_dim: int = 256,
) -> None:
    super().__init__()
    total_dim = geo_dim + scale_dim + viewing_dim + quality_dim
    self.net = nn.Sequential(
        nn.LayerNorm(total_dim),
        nn.Linear(total_dim, hidden_dim),
        nn.GELU(),
        nn.Linear(hidden_dim, embed_dim),
    )

MarsCLIPModel

MarsCLIPModel(*, vocab_size: int, image_size: int = 224, patch_size: int = 16, hidden_dim: int = 256, embed_dim: int = 256, text_max_length: int = 64, geo_dim: int = 8, scale_dim: int = 8, viewing_dim: int = 13, quality_dim: int = 7, mask_ratio: float = 0.0)

Bases: Module

First-pass tri-modal contrastive model for MarsCLIP.

Source code in src/clip/marsclip_model.py
def __init__(
    self,
    *,
    vocab_size: int,
    image_size: int = 224,
    patch_size: int = 16,
    hidden_dim: int = 256,
    embed_dim: int = 256,
    text_max_length: int = 64,
    geo_dim: int = 8,
    scale_dim: int = 8,
    viewing_dim: int = 13,
    quality_dim: int = 7,
    mask_ratio: float = 0.0,
) -> None:
    super().__init__()
    self.mask_ratio = mask_ratio
    self.image_tower = ImageViTTower(
        image_size=image_size,
        patch_size=patch_size,
        hidden_dim=hidden_dim,
        embed_dim=embed_dim,
    )
    self.text_tower = TextTransformerTower(
        vocab_size=vocab_size,
        max_length=text_max_length,
        hidden_dim=hidden_dim,
        embed_dim=embed_dim,
    )
    self.geo_tower = GeoContextTower(
        geo_dim=geo_dim,
        scale_dim=scale_dim,
        viewing_dim=viewing_dim,
        quality_dim=quality_dim,
        hidden_dim=hidden_dim,
        embed_dim=embed_dim,
    )
    self.logit_scale = nn.Parameter(torch.tensor(math.log(1 / 0.07), dtype=torch.float32))

compute_patch_valid_mask

compute_patch_valid_mask(valid_mask: Tensor, patch_size: int) -> torch.Tensor

Aggregate a pixel valid mask into one boolean per image patch.

Source code in src/clip/marsclip_model.py
def compute_patch_valid_mask(valid_mask: torch.Tensor, patch_size: int) -> torch.Tensor:
    """Aggregate a pixel valid mask into one boolean per image patch."""
    if valid_mask.ndim != 3:
        raise ValueError("valid_mask must have shape (B, H, W)")
    b, h, w = valid_mask.shape
    if h % patch_size != 0 or w % patch_size != 0:
        raise ValueError("Image size must be divisible by patch_size.")

    patches = valid_mask.unfold(1, patch_size, patch_size).unfold(2, patch_size, patch_size)
    return patches.any(dim=-1).any(dim=-1).reshape(b, -1)

sample_patch_keep_mask

sample_patch_keep_mask(patch_valid_mask: Tensor, mask_ratio: float, generator: Generator | None = None) -> torch.Tensor

Randomly keep a subset of valid image patches, never selecting invalid ones.

Source code in src/clip/marsclip_model.py
def sample_patch_keep_mask(
    patch_valid_mask: torch.Tensor,
    mask_ratio: float,
    generator: torch.Generator | None = None,
) -> torch.Tensor:
    """Randomly keep a subset of valid image patches, never selecting invalid ones."""
    if not (0.0 <= mask_ratio < 1.0):
        raise ValueError("mask_ratio must satisfy 0 <= mask_ratio < 1.")

    keep_mask = torch.zeros_like(patch_valid_mask, dtype=torch.bool)
    for i in range(patch_valid_mask.shape[0]):
        valid_idx = torch.nonzero(patch_valid_mask[i], as_tuple=False).flatten()
        if len(valid_idx) == 0:
            continue
        n_keep = max(1, int(math.ceil(len(valid_idx) * (1.0 - mask_ratio))))
        perm = torch.randperm(len(valid_idx), generator=generator, device=valid_idx.device)
        selected = valid_idx[perm[:n_keep]]
        keep_mask[i, selected] = True
    return keep_mask

symmetric_contrastive_loss

symmetric_contrastive_loss(a: Tensor, b: Tensor, logit_scale: Tensor) -> torch.Tensor

CLIP-style symmetric contrastive loss for one embedding pair.

Source code in src/clip/marsclip_model.py
def symmetric_contrastive_loss(
    a: torch.Tensor,
    b: torch.Tensor,
    logit_scale: torch.Tensor,
) -> torch.Tensor:
    """CLIP-style symmetric contrastive loss for one embedding pair."""
    logits = torch.matmul(a, b.T) * logit_scale.exp()
    targets = torch.arange(a.shape[0], device=a.device)
    return 0.5 * (
        F.cross_entropy(logits, targets) + F.cross_entropy(logits.T, targets)
    )