Skip to content

Commit 6ec67fe

Browse files
author
Peter Manshausen
committed
Add script to convert EMA network snapshots to inference checkpoints
Training writes EMA weights only to network-snapshot-*.pkl. This script wraps the 'ema' module with batch_info and model_config from a training-state checkpoint of the same run, so the EMA weights can be loaded with cbottle.inference.load. Signed-off-by: Peter Manshausen <pmanshausen@nvidia.com>
1 parent 4ce9aee commit 6ec67fe

1 file changed

Lines changed: 87 additions & 0 deletions

File tree

Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
#!/usr/bin/env python3
2+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
3+
# SPDX-License-Identifier: Apache-2.0
4+
#
5+
# Licensed under the Apache License, Version 2.0 (the "License");
6+
# you may not use this file except in compliance with the License.
7+
# You may obtain a copy of the License at
8+
#
9+
# http://www.apache.org/licenses/LICENSE-2.0
10+
#
11+
# Unless required by applicable law or agreed to in writing, software
12+
# distributed under the License is distributed on an "AS IS" BASIS,
13+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14+
# See the License for the specific language governing permissions and
15+
# limitations under the License.
16+
17+
import argparse
18+
import pickle
19+
from pathlib import Path
20+
21+
from cbottle import checkpointing
22+
23+
24+
def get_args():
25+
parser = argparse.ArgumentParser(
26+
description=(
27+
"Convert a network-snapshot-*.pkl EMA artifact into a "
28+
"training-state-style .checkpoint file for inference."
29+
)
30+
)
31+
parser.add_argument(
32+
"--snapshot",
33+
type=str,
34+
required=True,
35+
help="Path to network-snapshot-*.pkl",
36+
)
37+
parser.add_argument(
38+
"--reference-checkpoint",
39+
type=str,
40+
required=True,
41+
help=(
42+
"Path to a training-state-*.checkpoint from the same run. "
43+
"Used to copy batch_info and model_config."
44+
),
45+
)
46+
parser.add_argument(
47+
"--output",
48+
type=str,
49+
required=True,
50+
help="Path to write converted EMA checkpoint.",
51+
)
52+
return parser.parse_args()
53+
54+
55+
def main():
56+
args = get_args()
57+
58+
snapshot_path = Path(args.snapshot).expanduser().resolve()
59+
reference_ckpt_path = Path(args.reference_checkpoint).expanduser().resolve()
60+
output_path = Path(args.output).expanduser().resolve()
61+
62+
if not snapshot_path.is_file():
63+
raise FileNotFoundError(str(snapshot_path))
64+
if not reference_ckpt_path.is_file():
65+
raise FileNotFoundError(str(reference_ckpt_path))
66+
67+
with open(snapshot_path, "rb") as f:
68+
snapshot = pickle.load(f)
69+
if "ema" not in snapshot:
70+
raise KeyError(f"{snapshot_path} does not contain key 'ema'")
71+
ema_model = snapshot["ema"]
72+
73+
with checkpointing.Checkpoint(str(reference_ckpt_path), "r") as ref_ckpt:
74+
batch_info = ref_ckpt.read_batch_info()
75+
model_config = ref_ckpt.read_model_config()
76+
77+
output_path.parent.mkdir(parents=True, exist_ok=True)
78+
with checkpointing.Checkpoint(str(output_path), "w") as out_ckpt:
79+
out_ckpt.write_model(ema_model)
80+
out_ckpt.write_batch_info(batch_info)
81+
out_ckpt.write_model_config(model_config)
82+
83+
print(f"Wrote EMA checkpoint: {output_path}")
84+
85+
86+
if __name__ == "__main__":
87+
main()

0 commit comments

Comments
 (0)