xdecoder/Instruct-X-Decoder
163
1import pickle2from distutils import log3 4import torch5import torch.nn.functional as F6import torch.distributed as dist7 8from einops import rearrange, repeat9from timm.loss import SoftTargetCrossEntropy10 11soft_cross_entropy = SoftTargetCrossEntropy()12 13def is_dist_initialized():14 return torch.distributed.is_initialized()15 16def get_world_size():17 if is_dist_initialized():18 return torch.distributed.get_world_size()19 return 120 21def get_rank():22 if is_dist_initialized():23 return dist.get_rank()24 return 025 26def all_gather_grad(x):27 if get_world_size() > 1:28 all_x = [torch.zeros_like(x) for _ in range(get_world_size())]29 torch.distributed.all_gather(all_x, x)30 all_x[torch.distributed.get_rank()] = x31 x = torch.cat(all_x, dim=0)32 return x33 34def vl_multilabel_contrastive_loss(image_feat, text_feat, temperature=1):35 """36 Args:37 image_feat (torch.Tensor): shape [B, L1, C] # B: batch_size, L1: 1, C: 25638 text_feat (torch.Tensor): shape [B, L2, C] # B:batch_size, L2: number of selected nouns, C: 25639 40 Returns:41 """42 # [B, L1, C], L1 = 143 # image_feat = F.normalize(image_feat, dim=-1)44 # [B, L2, C]45 # text_feat = F.normalize(text_feat, dim=-1)46 # HACK: normalize outside47 48 # [B, L1, L2]49 dist_per_img = image_feat @ rearrange(text_feat, 'b l c -> b c l') 50 # [B, L2, L1]51 dist_per_text = text_feat @ rearrange(image_feat, 'b l c -> b c l')52 53 batch = image_feat.shape[0]54 img_len = image_feat.shape[1]55 text_len = text_feat.shape[1]56 # [B, L1, L2]57 pos_labels_batch_img = rearrange(torch.ones_like(dist_per_text) / dist_per_text.size(1), 'b l2 l1 -> b l1 l2')58 # [B, L2, L1]59 pos_labels_batch_text = rearrange(torch.ones_like(dist_per_img) / dist_per_img.size(1), 'b l1 l2 -> b l2 l1')60 61 image_x = rearrange(image_feat, 'b l c -> (b l) c')62 text_x = rearrange(text_feat, 'b l c -> (b l) c')63 64 logits_per_img = image_x @ all_gather_grad(text_x).t()65 logits_per_text = text_x @ all_gather_grad(image_x).t()66 67 # get label globally68 # [B, L1, B, L2, W]69 labels_per_img = F.one_hot(70 torch.ones(batch, img_len, batch, text_len, dtype=torch.long, device=image_x.device) * get_rank(),71 num_classes=get_world_size()).to(image_x.dtype)72 labels_per_img *= rearrange(pos_labels_batch_img, 'b l1 l2 -> b l1 1 l2 1') * repeat(73 torch.eye(batch, dtype=image_x.dtype, device=image_x.device), 'b1 b2 -> b1 1 b2 1 1')74 # [BxL1, WxBxL2]75 labels_per_img = rearrange(labels_per_img, 'b1 l1 b2 l2 w -> (b1 l1) (w b2 l2)')76 # [B, L2, B, L1, W]77 labels_per_text = F.one_hot(78 torch.ones(batch, text_len, batch, img_len, dtype=torch.long, device=text_x.device) * get_rank(),79 num_classes=get_world_size()).to(text_x.dtype)80 labels_per_text *= rearrange(pos_labels_batch_text, 'b l2 l1 -> b l2 1 l1 1') * repeat(81 torch.eye(batch, dtype=text_x.dtype, device=image_x.device), 'b2 b1 -> b2 1 b1 1 1')82 # [BxL2, WxBxL1]83 labels_per_text = rearrange(labels_per_text, 'b2 l2 b1 l1 w -> (b2 l2) (w b1 l1)')84 85 logit_scale = temperature.exp().clamp(max=100)86 87 loss_img = soft_cross_entropy(logit_scale * logits_per_img, labels_per_img)88 loss_text = soft_cross_entropy(logit_scale * logits_per_text, labels_per_text)89 90 loss = 0.5 * (loss_img + loss_text)91 return loss92 93def vl_contrastive_loss(image_feat, text_feat, temperature=1):94 # if image_id or text_id is None, it should be None across all GPUs95 # image_feat = F.normalize(image_feat, dim=1)96 # text_feat = F.normalize(text_feat, dim=1)97 # handle normalization outside98 99 # add the following 4 lines100 image_feat = all_gather_grad(image_feat)101 text_feat = all_gather_grad(text_feat)102 103 logits = torch.matmul(image_feat, text_feat.t())104 logit_scale = temperature.exp().clamp(max=100)105 106 gt = torch.arange(logits.shape[0], device=logits.device)107 loss1 = F.cross_entropy(logit_scale * logits, gt)108 loss2 = F.cross_entropy(logit_scale * logits.t(), gt)109 return (loss1 + loss2) / 2 # scale it up by the number of GPUs110 111 112def all_gather_pickle(data, device):113 """114 Run all_gather on arbitrary picklable data (not necessarily tensors)115 Args:116 data: any picklable object117 Returns:118 list[data]: list of data gathered from each rank119 """120 world_size = get_world_size()121 if world_size == 1:122 return [data]123 124 # serialized to a Tensor125 buffer = pickle.dumps(data)126 storage = torch.ByteStorage.from_buffer(buffer)127 tensor = torch.ByteTensor(storage).to(device)128 129 # obtain Tensor size of each rank130 local_size = torch.LongTensor([tensor.numel()]).cuda()131 size_list = [torch.LongTensor([0]).cuda() for _ in range(world_size)]132 dist.all_gather(size_list, local_size)133 size_list = [int(size.item()) for size in size_list]134 max_size = max(size_list)135 136 # receiving Tensor from all ranks137 # we pad the tensor because torch all_gather does not support138 # gathering tensors of different shapes139 tensor_list = []140 for _ in size_list:141 tensor_list.append(torch.ByteTensor(size=(max_size,)).cuda())142 if local_size != max_size:143 padding = torch.ByteTensor(size=(max_size - local_size,)).cuda()144 tensor = torch.cat((tensor, padding), dim=0)145 dist.all_gather(tensor_list, tensor)146 147 data_list = []148 for size, tensor in zip(size_list, tensor_list):149 buffer = tensor.cpu().numpy().tobytes()[:size]150 data_list.append(pickle.loads(buffer))151 152 return data_list153 154def all_gather_arbitary_tensor(tensor):155 if get_world_size() > 1:156 device = tensor.device157 tensor_batch = all_gather_pickle(tensor.cpu(), device)158 tensor_batch = [x.to(device) for x in tensor_batch]159 tensor_batch[torch.distributed.get_rank()] = tensor160 tensor_batch = torch.cat(tensor_batch, dim=0)161 else:162 tensor_batch = tensor163 return tensor_batch164 165def ql_contrastive_loss(image_feat, text_feat, temperature=1):166 # add the following 4 lines167 image_feat = all_gather_arbitary_tensor(image_feat)168 text_feat = all_gather_arbitary_tensor(text_feat)169 170 logits = torch.matmul(image_feat, text_feat.t())171 logit_scale = temperature.exp().clamp(max=100)172 173 gt = torch.arange(logits.shape[0], device=logits.device)174 loss1 = F.cross_entropy(logit_scale * logits, gt)175 loss2 = F.cross_entropy(logit_scale * logits.t(), gt)176 return (loss1 + loss2) / 2 # scale it up by the number of GPUs177 178def vl_similarity(image_feat, text_feat, temperature=1):179 # Only support single GPU for now.180 logits = torch.matmul(image_feat, text_feat.t())181 logits = temperature.exp().clamp(max=100) * logits182 return logits183 184def ql_multi_contrastive_loss(image_feat, text_feat, text_hash, temperature=1):185 # add the following 4 lines186 image_feat = all_gather_arbitary_tensor(image_feat)187 text_feat = all_gather_arbitary_tensor(text_feat)188 189 text_hash_batch = all_gather_pickle(text_hash, text_feat.device)190 text_hash_all = torch.cat(text_hash_batch)191 192 text_hash_all_unique = torch.unique(text_hash_all).tolist()193 gt = torch.zeros((image_feat.shape[0], len(text_hash_all_unique)), device=text_feat.device)194 text_hash_all = text_hash_all.tolist()195 text_feat_unique = torch.stack([text_feat[text_hash_all.index(txt)] for txt in text_hash_all_unique])196 197 for idx, txt in enumerate(text_hash_all):198 gt[idx][text_hash_all_unique.index(txt)] = 1199 200 logits = torch.matmul(image_feat, text_feat_unique.t())201 logits = logits*temperature.exp().clamp(max=100)202 203 loss_img = soft_cross_entropy(logits, gt)204 loss_text = soft_cross_entropy(logits.t(), gt.t() / gt.t().sum(-1, keepdim=True))205 206 loss = 0.7 * loss_img + 0.3 * loss_text207 return loss208 209def image_text_contrastive_loss_queue(image_feat_inp, text_feat_inp, lang_enc, training):210 # add the following 4 lines211 image_feat = all_gather_grad(image_feat_inp.contiguous())212 text_feat = all_gather_grad(text_feat_inp.contiguous())213 214 image_feat = image_feat / (image_feat.norm(dim=-1, keepdim=True) + 1e-7)215 text_feat = text_feat / (text_feat.norm(dim=-1, keepdim=True) + 1e-7)216 217 temperature = lang_enc.logit_scale218 logits = torch.matmul(image_feat, text_feat.t())219 logit_scale = temperature.exp().clamp(max=100)220 221 gt = torch.arange(logits.shape[0], device=logits.device)222 loss1 = F.cross_entropy(logit_scale * logits, gt)223 loss2 = F.cross_entropy(logit_scale * logits.t(), gt)224 225 return (loss1 + loss2) / 2 # scale it up by the number of GPUs