# Copyright 2024 EPFL and Apple Inc. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. import hashlib import torch import collections.abc from itertools import repeat import torchvision.transforms.functional as TF from fourm.utils.data_constants import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD def image_mask(tensor: torch.Tensor, GT_tokens: int, input_budget: int, target_budget: int): """Applies input and target masking to an image tensor sequentially Args: tensor: Image tensor GT_tokens: Number of tokens in the tensor input_budget: Token budget for the input target_budget: Token budget for the target Returns: Dictionary containing the masked image tensor, the input mask, the target mask, and the decoder attention mask """ # Input mask: First `input_budget` tokens are not masked (0), rest are masked (1) input_mask = torch.ones(GT_tokens, dtype=torch.bool) input_mask[:input_budget] = 0 # First `input_budget` positions are not masked # Target mask: The next `target_budget` tokens are not masked (0), rest are masked (1) target_mask = torch.ones(GT_tokens, dtype=torch.bool) if target_budget is not None: target_mask[input_budget:input_budget + target_budget] = 0 # Next `target_budget` positions are not masked else: target_mask = ~input_mask # If target_budget is None, complement input_mask # Compute decoder attention mask decoder_attention_mask = torch.zeros(GT_tokens, dtype=torch.int) first_mask_token = torch.argmin(target_mask + torch.arange(target_mask.shape[0], device=target_mask.device) * 1e-6) decoder_attention_mask[first_mask_token] = (~target_mask).sum() # Equivalent to target budget return { "tensor": torch.tensor(tensor).long().cuda(), "input_mask": input_mask.unsqueeze(0).cuda(), "target_mask": target_mask.unsqueeze(0).cuda(), "decoder_attention_mask": decoder_attention_mask.unsqueeze(0).cuda(), } def denormalize(img, mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD): """ Denormalizes an image. Args: img (torch.Tensor): Image to denormalize. mean (tuple): Mean to use for denormalization. std (tuple): Standard deviation to use for denormalization. """ return TF.normalize( img.clone(), mean= [-m/s for m, s in zip(mean, std)], std= [1/s for s in std] ) def generate_uint15_hash(seed_str): """Generates a hash of the seed string as an unsigned int15 integer""" return int(hashlib.sha256(seed_str.encode('utf-8')).hexdigest(), 16) % (2**15) # From PyTorch internals def _ntuple(n): def parse(x): if isinstance(x, collections.abc.Iterable): return x return tuple(repeat(x, n)) return parse to_1tuple = _ntuple(1) to_2tuple = _ntuple(2) to_3tuple = _ntuple(3) to_4tuple = _ntuple(4) to_ntuple = _ntuple