MCPcopy Create free account
hub / github.com/tinygrad/tinygrad / split_safetensor

Function split_safetensor

examples/webgpu/stable_diffusion/compile.py:32–71  ·  view source on GitHub ↗
(fn)

Source from the content-addressed store, hash-verified

30 rest_float32_values.tofile(f)
31
32def split_safetensor(fn):
33 _, data_start, metadata = safe_load_metadata(fn)
34 text_model_offset = 3772703308
35 chunk_size = 536870912
36
37 for k in metadata:
38 # safetensor is in fp16, except for text moel
39 if (metadata[k]["data_offsets"][0] < text_model_offset):
40 metadata[k]["data_offsets"][0] = int(metadata[k]["data_offsets"][0]/2)
41 metadata[k]["data_offsets"][1] = int(metadata[k]["data_offsets"][1]/2)
42
43 last_offset = 0
44 part_end_offsets = []
45
46 for k in metadata:
47 offset = metadata[k]['data_offsets'][0]
48
49 if offset == text_model_offset:
50 break
51
52 part_offset = offset - last_offset
53
54 if (part_offset >= chunk_size):
55 part_end_offsets.append(data_start+offset)
56 last_offset = offset
57
58 text_model_start = int(text_model_offset/2)
59 net_bytes = bytes(open(fn, 'rb').read())
60 part_end_offsets.append(text_model_start+data_start)
61 cur_pos = 0
62
63 for i, end_pos in enumerate(part_end_offsets):
64 with open(os.path.join(os.path.dirname(__file__), f'./net_part{i}.safetensors'), "wb+") as f:
65 f.write(net_bytes[cur_pos:end_pos])
66 cur_pos = end_pos
67
68 with open(os.path.join(os.path.dirname(__file__), f'./net_textmodel.safetensors'), "wb+") as f:
69 f.write(net_bytes[text_model_start+data_start:])
70
71 return part_end_offsets
72
73def fetch_dep(file, url):
74 with open(file, "w", encoding="utf-8") as f:

Callers 1

compile.pyFile · 0.85

Calls 4

safe_load_metadataFunction · 0.90
appendMethod · 0.80
readMethod · 0.45
writeMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…