MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT-LLM / deserialize_managed_weights

Function deserialize_managed_weights

tensorrt_llm/builder.py:938–969  ·  view source on GitHub ↗
(path: str | Path)

Source from the content-addressed store, hash-verified

936
937
938def 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
972def build(model: PretrainedModel, build_config: BuildConfig) -> Engine:

Callers 1

from_dirMethod · 0.85

Calls 1

decodeMethod · 0.45

Tested by

no test coverage detected