-
-
Save Narsil/3edeec2669a5e94e4707aa0f901d2282 to your computer and use it in GitHub Desktop.
| import mmap | |
| import torch | |
| import json | |
| import os | |
| from huggingface_hub import hf_hub_download | |
| def load_file(filename, device): | |
| with open(filename, mode="r", encoding="utf8") as file_obj: | |
| with mmap.mmap(file_obj.fileno(), length=0, access=mmap.ACCESS_READ) as m: | |
| header = m.read(8) | |
| n = int.from_bytes(header, "little") | |
| metadata_bytes = m.read(n) | |
| metadata = json.loads(metadata_bytes) | |
| size = os.stat(filename).st_size | |
| storage = torch.ByteStorage.from_file(filename, shared=False, size=size).untyped() | |
| offset = n + 8 | |
| return {name: create_tensor(storage, info, offset) for name, info in metadata.items() if name != "__metadata__"} | |
| DTYPES = {"F32": torch.float32} | |
| device = "cpu" | |
| def create_tensor(storage, info, offset): | |
| dtype = DTYPES[info["dtype"]] | |
| shape = info["shape"] | |
| start, stop = info["data_offsets"] | |
| return torch.asarray(storage[start + offset : stop + offset], dtype=torch.uint8).view(dtype=dtype).reshape(shape) | |
| def main(): | |
| filename = hf_hub_download("gpt2", filename="model.safetensors") | |
| weights = load_file(filename, device) | |
| print(weights.keys()) | |
| if __name__ == "__main__": | |
| main() |
When will torch get default safetesnors support?
you should open that request in torch repo, i think it'd be awesome
Great, will do that.
python # Preview a . safetensors file using pure PyTorch-only style loading logic. # Note: actual safetensors parsing requires the safetensors package or a custom parser. Import torch Path = "model . safetensors" #Placeholder preview: # 1) Inspect file metadate if available # 2) Load tensors into a state_dict - like structure # 3 ) Print tensor names, shapes, and dtypes Print ( f "Priviewing: {path} " ) compatible loader to read tesor keys and metadata . " ) print ( "Then compare against your model ' s expected state_dict keys . " ) # Examle of what you ' d print after loading : # for name, tensor in state_direct . items ( ) : # print ( name , tensor . shape , tensor . dtype)
Seems
deviceis not being used here.