Download GeoText-1652/Method/accelerators/accelerator.py from geobase/GeoText1652_model: direct link, hf CLI and curl.
- Browser
- Download file 1.06 kB
-
https://huggingface.co/geobase/GeoText1652_model/resolve/main/GeoText-1652/Method/accelerators/accelerator.py
- Command line
-
hf download hf://geobase/GeoText1652_model/GeoText-1652/Method/accelerators/accelerator.py
-
curl -L -o accelerator.py https://huggingface.co/geobase/GeoText1652_model/resolve/main/GeoText-1652/Method/accelerators/accelerator.py
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) | |