Safetensors
mhassanch's picture
Add GeoText inference endpoint support
4d3c316
Raw History Blame Contribute Delete
1.06 kB
# -*- coding: utf-8 -*-
# Multi-Grained Vision Language Pre-Training: Aligning Texts with Visual Concepts (https://arxiv.org/abs/2111.08276)
# Github: https://github.com/zengyan-97/X-VLM
# Copyright (c) 2022, ByteDance Inc.
# All rights reserved.
from logging import Logger
import torch
from torch.optim import Optimizer
Net = torch.nn.Module
class Accelerator:
def __init__(self, cfg, logger) -> None:
self.cfg = cfg
self.logger = logger
def set_up(self, model: Net):
raise NotImplementedError("Set Up method not implement in Accelerator, please check! ")
def broadcast(self):
raise NotImplementedError("Broadcast method not implement in Accelerator, please check! ")
def backward_step(self, loss: torch.Tensor):
loss.backward()
def optimizer_step(self, optimizer: Optimizer, model: Net, grad_norm: float) -> float:
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(),
grad_norm)
return float(total_norm)