oxfrug's picture
load.py docstring
baf5c33 verified
Raw
History Blame Contribute Delete
2.2 kB
"""Load oxfrug/chronos-2-int8-torchao as a Chronos2Pipeline.
pip install 'chronos-forecasting>=2.0' torchao safetensors
from load import load
pipe = load('oxfrug/chronos-2-int8-torchao', device='cuda')
Do not use Chronos2Pipeline.from_pretrained on this repo — the packed
INT8 tensors need this loader.
"""
from pathlib import Path
import json, torch
from safetensors.torch import load_file
from torchao.quantization import Int8Tensor, Int8WeightOnlyConfig, quantize_
from transformers import AutoConfig
from chronos import Chronos2Pipeline
from chronos.chronos2.model import Chronos2Model
SKIP = ('output_patch_embedding', 'input_embed', 'shared')
def _skip(mod, name):
if not isinstance(mod, torch.nn.Linear): return False
if any(s in name.lower() for s in SKIP): return False
if mod.weight.numel() < 4096: return False
return True
def load(repo_or_dir, device='cuda'):
from huggingface_hub import snapshot_download
p = Path(repo_or_dir)
if not (p / 'model.safetensors').exists():
p = Path(snapshot_download(repo_or_dir))
meta = json.loads((p / 'quant_meta.json').read_text())
cfg = AutoConfig.from_pretrained(p)
model = Chronos2Model(cfg)
quantize_(model, Int8WeightOnlyConfig(), filter_fn=_skip)
flat = load_file(str(p / 'model.safetensors'))
dense = {k[7:]: t for k, t in flat.items() if k.startswith('dense::')}
sd = model.state_dict()
for k, t in dense.items():
if k in sd and sd[k].shape == t.shape:
sd[k].copy_(t.to(sd[k].device))
for full_name, spec in meta['int8'].items():
prefix = f'i8::{full_name}::'
obj = model
parts = full_name.split('.')
for part in parts[:-1]:
obj = getattr(obj, part)
packed = Int8Tensor(flat[prefix+'qdata'], flat[prefix+'scale'], spec['block_size'], torch.float32, zero_point=flat[prefix+'zero_point'])
if parts[-1] == 'weight' and isinstance(obj, torch.nn.Linear):
obj.weight = torch.nn.Parameter(packed, requires_grad=False)
else:
setattr(obj, parts[-1], packed)
if device == 'cuda':
model = model.cuda()
return Chronos2Pipeline(model=model)