Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
loss.py225 linesDownload Raw Back to language
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