| 936 | |
| 937 | |
| 938 | def deserialize_managed_weights(path: str | Path) -> dict[str, np.ndarray]: |
| 939 | with open(path, "rb") as f: |
| 940 | header_json_len = int.from_bytes(f.read(8), byteorder="little") |
| 941 | header_json = f.read(header_json_len).decode() |
| 942 | header = json.loads(header_json) |
| 943 | |
| 944 | managed_weights = {} |
| 945 | for name, info in header.items(): |
| 946 | dtype = info["dtype"] |
| 947 | shape = info["shape"] |
| 948 | data_offsets = info["data_offsets"] |
| 949 | if dtype == "F32": |
| 950 | dtype = np.float32 |
| 951 | elif dtype == "F16": |
| 952 | dtype = np.float16 |
| 953 | elif dtype == "BF16": |
| 954 | dtype = np_bfloat16 |
| 955 | elif dtype == "F8_E4M3": |
| 956 | dtype = np_float8 |
| 957 | elif dtype == "I64": |
| 958 | dtype = np.int64 |
| 959 | elif dtype == "I32": |
| 960 | dtype = np.int32 |
| 961 | else: |
| 962 | raise RuntimeError(f"Unsupported dtype: {dtype}") |
| 963 | |
| 964 | f.seek(data_offsets[0] + header_json_len + 8) |
| 965 | buf = f.read(data_offsets[1] - data_offsets[0]) |
| 966 | value = np.frombuffer(buf, dtype=dtype).reshape(shape) |
| 967 | managed_weights[name] = value |
| 968 | |
| 969 | return managed_weights |
| 970 | |
| 971 | |
| 972 | def build(model: PretrainedModel, build_config: BuildConfig) -> Engine: |