zry-research commited on
Commit
d0338a5
·
1 Parent(s): 8ca2928

feat: add weight to model/

Browse files
inference.py ADDED
@@ -0,0 +1,439 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # -*- coding: utf-8 -*-
2
+ """
3
+ VLLM Information Density & Visual Comfort Auditor
4
+ 功能:批量扫描文件夹内的图片,基于"Visual Comfort"标准判断图片是 Suitable 还是 Unsuitable。
5
+ 输出:终端统计报告 + 详细 JSON 结果文件
6
+
7
+ 使用方法:
8
+ python rule_info_vllm.py \
9
+ --input_dir "/path/to/your/images" \
10
+ --model_path "/path/to/your/Qwen2.5-VL-7B-Instruct"
11
+ """
12
+
13
+ import os
14
+ import re
15
+ import json
16
+ import argparse
17
+ import multiprocessing
18
+ from pathlib import Path
19
+ from typing import Dict, List, Optional, Any
20
+ from tqdm import tqdm
21
+ from PIL import Image, ImageFile
22
+ from transformers import AutoProcessor
23
+ from vllm import LLM, SamplingParams
24
+
25
+ # ==========================================
26
+ # 【核心配置】环境与并发设置
27
+ # ==========================================
28
+ os.environ['VLLM_WORKER_MULTIPROC_METHOD'] = 'spawn'
29
+
30
+ # 防止部分图片因截断而报错
31
+ ImageFile.LOAD_TRUNCATED_IMAGES = True
32
+
33
+ # 支持的图片格式
34
+ IMG_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp", ".tif", ".tiff"}
35
+
36
+ # ==========================================
37
+ # 【提示词工程】
38
+ # ==========================================
39
+ SYS_PROMPT_TEXT = """ """
40
+
41
+
42
+ ##用法示例:
43
+ #SYS_PROMPT_TEXT 替换!
44
+
45
+ # --- 1. EXQUISITENESS (精美度) ---
46
+ EXQUISITENESS_SYSTEM_PROMPT = """You are a highly critical Senior Art Director and Visual Auditor.
47
+ Your task is to identify "Low-Quality, Amateur, or Overly Simplistic" advertising materials based on the "Exquisiteness" standard.
48
+ You have ZERO TOLERANCE for "Cheap Templates" that lack professional depth and aesthetic cohesion.
49
+
50
+ INPUT: One image and one natural-language question about aesthetic exquisiteness.
51
+
52
+ YOUR TASK:
53
+ 1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
54
+ 2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
55
+
56
+ OUTPUT FORMAT:
57
+ Return EXACTLY two blocks, no extra text:
58
+ <think>Detailed reasoning comparing visual features against BOTH unsuitable and suitable criteria (checking for template-like flatness vs. rich visual layers)...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Exquisiteness"}</answer>
59
+
60
+ =========================================
61
+ CRITERIA FOR 'UNSUITABLE' (VIOLATION / LOW QUALITY)
62
+ =========================================
63
+ 1. **Simplistic & "Flat" Design (The "Template" Trap):**
64
+ - **Overly Basic:** The layout is overly basic: just a simple solid color block + a generic icon + plain text.
65
+ - **Lack of Depth:** The design feels like a "default" or "low-end" template with zero artistic polish, shadows, or texture.
66
+ - **Disconnected:** Elements feel placed randomly without visual cohesion.
67
+
68
+ 2. **Poor Overall Aesthetics:**
69
+ - **Low Fidelity:** The colors are muddy, the composition is unbalanced, or the material quality looks pixelated/unrefined.
70
+ - **Cheap Experience:** The image fails to provide a "premium" or "high-fidelity" visual experience.
71
+
72
+ =========================================
73
+ CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
74
+ =========================================
75
+ 1. **Rich Visual Layers:** Use of depth, professional lighting, shadows, and high-quality textures.
76
+ 2. **Professional Polish:** The image features a coherent color palette and clear visual hierarchy that feels "designed" rather than "assembled".
77
+ 3. **Intentional Minimalism:** Even if the design is simple, it looks intentional, high-end, and balanced (not empty or basic).
78
+
79
+ =========================================
80
+ DECISION LOGIC
81
+ =========================================
82
+ - **Unsuitable**: If the design is simplistic, flat, low-effort, or looks like a cheap template.
83
+ - **Suitable**: If the image features rich visual layers, depth, and looks polished/premium.
84
+ """
85
+
86
+ # --- 2. PROFESSIONAL POLISH (后期质感) ---
87
+ PROFESSIONAL_POLISH_SYSTEM_PROMPT = """You are a highly critical Senior Art Director specializing in Post-Production.
88
+ Your task is to evaluate "Post-Production Quality" for S-level splash ads.
89
+ You have ZERO TOLERANCE for raw, unprocessed photos that look like amateur snapshots.
90
+
91
+ INPUT: One image and one natural-language question about post-production quality.
92
+
93
+ YOUR TASK:
94
+ 1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
95
+ 2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
96
+
97
+ OUTPUT FORMAT:
98
+ Return EXACTLY two blocks, no extra text:
99
+ <think>Detailed reasoning comparing visual features against BOTH unsuitable and suitable criteria (analyzing lighting, color grading, depth of field, and texture)...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Professional Polish"}</answer>
100
+
101
+ =========================================
102
+ CRITERIA FOR 'UNSUITABLE' (VIOLATION / AMATEUR SNAPSHOT)
103
+ =========================================
104
+ 1. **Lack of Professional Post-Processing (Raw Photo Look):**
105
+ - **Unprocessed:** The image appears to be a "Raw Photo" directly from a camera/phone without professional retouching.
106
+ - **Poor Technicals:** There is no deliberate optimization of Lighting (flat or messy light), Color (dull or unbalanced tones), or Depth of Field (lack of professional bokeh or focus control).
107
+
108
+ 2. **The "Amateur Snapshot" Aesthetic:**
109
+ - **Casual Vibe:** The image looks like something a "passerby" could easily capture. It lacks the sophisticated framing, high-end texture, and artistic polish required for premium advertising.
110
+ - **Cheap Feel:** The visual quality feels "Cheap" and fails to convey the premium value of the brand.
111
+
112
+ 3. **Absence of Value Conveyance:**
113
+ - **Visually Flat:** The image fails to evoke a sense of high quality or luxury. It is visually "flat" and does not use post-production techniques to guide emotions.
114
+
115
+ =========================================
116
+ CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
117
+ =========================================
118
+ 1. **Professional Polish:** The image has distinct, high-end color grading and lighting that creates a "Cinematic" or "Commercial" look.
119
+ 2. **Intentional Aesthetic:** Clear mastery of lighting (highlights/shadows), color harmony, and depth-of-field that elevates the subject.
120
+ 3. **Premium Value:** The overall visual execution feels expensive and exclusive, far beyond a casual photograph.
121
+ 4. **Cinematic Excellence:** Film stills or cinematic stills are ALWAYS classified as SUITABLE.
122
+
123
+ =========================================
124
+ DECISION LOGIC
125
+ =========================================
126
+ - **Unsuitable**: If the image looks like an unprocessed, amateur, or casual snapshot (Raw Photo).
127
+ - **Suitable**: If the image shows professional polish, cinematic lighting, or is a film still.
128
+ """
129
+
130
+ # --- 3. LAYOUT BREATHABILITY (布局呼吸感 - 已包含正向标准) ---
131
+ LAYOUT_BREATHABILITY_SYSTEM_PROMPT = """You are a highly critical Senior Art Director specializing in Layout and Visual Hierarchy.
132
+ Your task is to identify "Suffocating Designs"—creative pieces where elements are too cramped, lack breathing room, or feel disorganized due to poor spacing.
133
+
134
+ INPUT: One image and one natural-language question about layout spacing.
135
+
136
+ YOUR TASK:
137
+ 1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
138
+ 2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
139
+
140
+ OUTPUT FORMAT:
141
+ Return EXACTLY two blocks, no extra text:
142
+ <think>Detailed reasoning comparing visual features against BOTH unsuitable and suitable criteria (analyzing module gap, edge tension, and visual path)...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Layout Breathability Check"}</answer>
143
+
144
+ =========================================
145
+ CRITERIA FOR 'UNSUITABLE' (VIOLATION / SUFFOCATING DESIGN)
146
+ =========================================
147
+ 1. **Lack of Breathing Room (Crowded Modules):**
148
+ - **Core Violation:** The main subject (product/hero), the headline text, and the logo are placed too close to each other.
149
+ - **Visual Feel:** The overall design feels "heavy" or "claustrophobic" because major modules lack sufficient negative space.
150
+ - **Small Print Nuance:** While secondary text (annotations) can have smaller gaps, they must NOT feel like they are "clinging" or "tangent" to other elements.
151
+
152
+ 2. **Edge Tension (贴边风险):**
153
+ - **Tangency:** Elements are unintentionally "touching" or "tangent" to each other or the canvas border without intentional overlapping (creating uncomfortable tension).
154
+
155
+ 3. **Information Overload:**
156
+ - **Clutter:** The layout is filled with too many text blocks or icons with no clear separation.
157
+ - **No Visual Path:** The eye doesn't know where to rest because every element is competing for attention and space simultaneously.
158
+
159
+ =========================================
160
+ CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
161
+ =========================================
162
+ 1. **Generous White Space:** Clear and deliberate separation between the headline, the hero subject, and the footer information.
163
+ 2. **Structured Layout:** Elements follow a clear grid or intentional alignment that allows the design to "breathe" while maintaining a strong hierarchy.
164
+ 3. **Intentional Overlap:** If elements overlap, it looks artistic, layered, and deliberate (e.g., text weaving behind a subject), not accidental or messy.
165
+
166
+ =========================================
167
+ DECISION LOGIC
168
+ =========================================
169
+ - **Unsuitable**: If the layout feels squeezed, crowded, has uncomfortable edge tension, or lacks a visual path.
170
+ - **Suitable**: If the layout has generous negative space, structured alignment, and allows the eye to travel comfortably.
171
+ """
172
+
173
+ # --- 4. TEXT LEGIBILITY (文字易读性与排布) ---
174
+ TEXT_LEGIBILITY_SYSTEM_PROMPT = """You are a highly critical Senior Art Director.
175
+ Your task is to evaluate "Information Accessibility" to ensure advertising copy is instantly readable and strategically placed.
176
+
177
+ INPUT: One image and one natural-language question about text legibility.
178
+
179
+ YOUR TASK:
180
+ 1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
181
+ 2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
182
+
183
+ OUTPUT FORMAT:
184
+ Return EXACTLY two blocks, no extra text:
185
+ <think>Detailed reasoning comparing visual features against BOTH unsuitable and suitable criteria (analyzing contrast, background interference, and placement logic)...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Text Legibility and Placement"}</answer>
186
+
187
+ =========================================
188
+ CRITERIA FOR 'UNSUITABLE' (VIOLATION / POOR LEGIBILITY)
189
+ =========================================
190
+ 1. **Direct Overlay on Complex Background:**
191
+ - **Interference:** Text is placed over "busy" areas (faces, complex textures, high-contrast patterns) where background details "cut through" the strokes.
192
+ - **No Protection:** The text lacks professional treatments (no solid backing, masks, strokes, or drop shadows) to separate it from the noisy background.
193
+
194
+ 2. **Poor Visual Catchiness (Dead Zones):**
195
+ - **Placement:** The primary marketing message (the "Hook") is placed in a "visual dead zone" (extreme edges or corners) where the eye does not naturally land.
196
+ - **Weak Presence:** Important copy is too small or lacks contrast relative to its position, failing to be eye-catching.
197
+
198
+ 3. **Low Contrast / Visual Camouflage:**
199
+ - **Blending:** Text color is too similar to background colors, causing it to "camouflage".
200
+ - **No Hierarchy:** There is no clear visual hierarchy; the eye has to "search" or struggle to find the text.
201
+
202
+ 4. **Unclear Small Text:**
203
+ - **Illegible:** Footnotes, disclaimers, or annotations are buried in background noise and are difficult to read (excluding text naturally on product packaging).
204
+
205
+ =========================================
206
+ CRITERIA FOR 'SUITABLE' (NON-VIOLATION / HIGH ACCESSIBILITY)
207
+ =========================================
208
+ 1. **Clean Placement:** Main text is placed on a "clean" area of the image (e.g., sky, plain wall, or a blurred background) utilizing negative space.
209
+ 2. **Professional Treatment:** Even if the background is complex, the text uses solid containers, high-contrast strokes, masks, or heavy shadows to ensure perfect, instant legibility.
210
+ 3. **Strategic Positioning:** Key messages are placed in focal points, not hidden in corners.
211
+ 4. **Cinematic Exception:** Film stills or cinematic captures are ALWAYS classified as SUITABLE (Cinematic Excellence).
212
+
213
+ =========================================
214
+ DECISION LOGIC
215
+ =========================================
216
+ - **Unsuitable**: If text is hard to read due to low contrast, complex background interference, or placement in dead zones.
217
+ - **Suitable**: If text is instantly readable due to clean placement, professional contrast treatments (masks/shadows), or if it is a cinematic still.
218
+ """
219
+
220
+ # ==========================================
221
+ # 辅助函数
222
+ # ==========================================
223
+
224
+ def collect_images(input_dir: Path) -> List[Dict[str, str]]:
225
+ """扫描目录下所有图片"""
226
+ if not input_dir.exists():
227
+ raise FileNotFoundError(f"Input directory not found: {input_dir}")
228
+
229
+ files = [p for p in input_dir.iterdir() if p.is_file() and p.suffix.lower() in IMG_EXTS]
230
+ files.sort()
231
+
232
+ print(f"[Info] Found {len(files)} images in {input_dir}")
233
+ return [{"path": str(p), "filename": p.name} for p in files]
234
+
235
+ def parse_llm_output(text: str) -> Dict[str, Any]:
236
+ """
237
+ 解析模型输出,提取 <think> 和 <answer> 标签中的内容
238
+ 目标格式: <think>...</think><answer>{"Answer": "Suitable/Unsuitable", ...}</answer>
239
+ """
240
+ default_res = {
241
+ "label": "Parse Error",
242
+ "think": "No reasoning found",
243
+ "raw": text
244
+ }
245
+
246
+ if not text:
247
+ return default_res
248
+
249
+ # 1. 提取 <think> 内容
250
+ think_match = re.search(r'<think>(.*?)</think>', text, re.DOTALL)
251
+ think_content = think_match.group(1).strip() if think_match else ""
252
+
253
+ # 2. 提取 <answer> 内容
254
+ answer_match = re.search(r'<answer>(.*?)</answer>', text, re.DOTALL)
255
+
256
+ extracted_label = "Parse Error"
257
+
258
+ if answer_match:
259
+ json_str = answer_match.group(1).strip()
260
+ try:
261
+ # 尝试解析 JSON
262
+ data = json.loads(json_str)
263
+ # 获取 Answer 字段 (Suitable 或 Unsuitable)
264
+ raw_ans = data.get("Answer", "")
265
+
266
+ # 归一化处理
267
+ if "unsuitable" in raw_ans.lower():
268
+ extracted_label = "Unsuitable"
269
+ elif "suitable" in raw_ans.lower():
270
+ extracted_label = "Suitable"
271
+ else:
272
+ extracted_label = raw_ans # 如果是其他奇怪的内容,保留原样
273
+
274
+ except json.JSONDecodeError:
275
+ # 如果 JSON 解析失败,尝试暴力匹配字��串
276
+ if "Unsuitable" in json_str:
277
+ extracted_label = "Unsuitable"
278
+ elif "Suitable" in json_str:
279
+ extracted_label = "Suitable"
280
+ else:
281
+ # 兜底:如果没有 answer 标签,直接在全文搜
282
+ if "Unsuitable" in text:
283
+ extracted_label = "Unsuitable"
284
+ elif "Suitable" in text:
285
+ extracted_label = "Suitable"
286
+
287
+ return {
288
+ "label": extracted_label,
289
+ "think": think_content,
290
+ "raw": text
291
+ }
292
+
293
+ def prepare_vllm_inputs(batch_meta: List[Dict], processor) -> List[Dict]:
294
+ """构建 vLLM 输入格式"""
295
+ vllm_inputs = []
296
+ # 根据新的 Prompt 设置对应的用户提问
297
+ user_query = "Is this image visually comfortable and suitable for information display?"
298
+
299
+ for item in batch_meta:
300
+ img_path = item["path"]
301
+ try:
302
+ image_obj = Image.open(img_path).convert("RGB")
303
+
304
+ messages = [
305
+ {"role": "system", "content": [{"type": "text", "text": SYS_PROMPT_TEXT}]},
306
+ {"role": "user", "content": [
307
+ {"type": "image", "image": img_path},
308
+ {"type": "text", "text": user_query}
309
+ ]}
310
+ ]
311
+
312
+ prompt_text = processor.apply_chat_template(
313
+ messages, tokenize=False, add_generation_prompt=True
314
+ )
315
+
316
+ vllm_inputs.append({
317
+ "prompt": prompt_text,
318
+ "multi_modal_data": {"image": image_obj}
319
+ })
320
+ except Exception as e:
321
+ print(f"[Warning] Failed to load {img_path}: {e}")
322
+ vllm_inputs.append(None)
323
+
324
+ return vllm_inputs
325
+
326
+ # ==========================================
327
+ # 主程序
328
+ # ==========================================
329
+
330
+ def main():
331
+ parser = argparse.ArgumentParser(description="AI Visual Comfort Auditor")
332
+ parser.add_argument("--input_dir", type=str, required=True, help="Folder containing images to check")
333
+ parser.add_argument("--model_path", type=str, required=True, help="Path to local Qwen-VL model")
334
+ parser.add_argument("--batch_size", type=int, default=128, help="Inference batch size")
335
+ parser.add_argument("--tp_size", type=int, default=4, help="Tensor Parallel size")
336
+ args = parser.parse_args()
337
+
338
+ input_path = Path(args.input_dir)
339
+ meta_data = collect_images(input_path)
340
+
341
+ if not meta_data:
342
+ print("[Info] No images found. Exiting.")
343
+ return
344
+
345
+ # ---------------------------
346
+ # 初始化模型
347
+ # ---------------------------
348
+ print(f"\n[Init] Loading Model: {args.model_path}")
349
+
350
+ llm = LLM(
351
+ model=args.model_path,
352
+ tokenizer=args.model_path,
353
+ trust_remote_code=True,
354
+ tensor_parallel_size=args.tp_size,
355
+ gpu_memory_utilization=0.90,
356
+ max_model_len=8192,
357
+ enforce_eager=True,
358
+ limit_mm_per_prompt={"image": 1}
359
+ )
360
+
361
+ processor = AutoProcessor.from_pretrained(args.model_path, trust_remote_code=True)
362
+
363
+ # 采样参数
364
+ sampling_params = SamplingParams(
365
+ temperature=0.7, # 稍微降低温度以获得更稳定的分类
366
+ max_tokens=1024,
367
+ top_p=0.9
368
+ )
369
+
370
+ # ---------------------------
371
+ # 批量推理
372
+ # ---------------------------
373
+ results = []
374
+ print(f"\n[Run] Starting Inference on {len(meta_data)} images...")
375
+
376
+ for i in tqdm(range(0, len(meta_data), args.batch_size), desc="Processing Batches"):
377
+ batch_meta = meta_data[i : i + args.batch_size]
378
+ batch_inputs = prepare_vllm_inputs(batch_meta, processor)
379
+
380
+ valid_inputs = [inp for inp in batch_inputs if inp is not None]
381
+ valid_indices = [idx for idx, inp in enumerate(batch_inputs) if inp is not None]
382
+
383
+ if not valid_inputs:
384
+ continue
385
+
386
+ outputs = llm.generate(valid_inputs, sampling_params=sampling_params, use_tqdm=False)
387
+
388
+ for local_idx, out in enumerate(outputs):
389
+ original_meta = batch_meta[valid_indices[local_idx]]
390
+ generated_text = out.outputs[0].text
391
+
392
+ # 解析结果
393
+ parsed = parse_llm_output(generated_text)
394
+
395
+ results.append({
396
+ "filename": original_meta["filename"],
397
+ "path": original_meta["path"],
398
+ "label": parsed["label"], # Suitable / Unsuitable
399
+ "think": parsed["think"], # 思维链
400
+ "raw_output": generated_text
401
+ })
402
+
403
+ # ---------------------------
404
+ # 统计与输出
405
+ # ---------------------------
406
+ total = len(results)
407
+ unsuitable_count = sum(1 for r in results if r["label"] == "Unsuitable")
408
+ suitable_count = sum(1 for r in results if r["label"] == "Suitable")
409
+ error_count = total - unsuitable_count - suitable_count
410
+
411
+ unsuitable_rate = (unsuitable_count / total * 100) if total > 0 else 0
412
+ suitable_rate = (suitable_count / total * 100) if total > 0 else 0
413
+
414
+ print("\n" + "="*60)
415
+ print(f"AUDIT REPORT FOR: {input_path.name}")
416
+ print("="*60)
417
+ print(f"{'Total Images':<25}: {total}")
418
+ print("-" * 60)
419
+ print(f"{'UNSUITABLE (Violation)':<25}: {unsuitable_count} ({unsuitable_rate:.2f}%)")
420
+ print(f"{'SUITABLE (Safe)':<25}: {suitable_count} ({suitable_rate:.2f}%)")
421
+ print(f"{'Parse Errors':<25}: {error_count}")
422
+ print("="*60)
423
+
424
+ # 保存结果
425
+ output_file = input_path / f"audit_result_{input_path.name}.json"
426
+ try:
427
+ with open(output_file, "w", encoding="utf-8") as f:
428
+ json.dump(results, f, ensure_ascii=False, indent=2)
429
+ print(f"\n[Done] Detailed JSON report saved to:\n-> {output_file}")
430
+ except Exception as e:
431
+ print(f"[Error] Could not save JSON: {e}")
432
+
433
+ if __name__ == "__main__":
434
+ try:
435
+ multiprocessing.set_start_method('spawn', force=True)
436
+ except RuntimeError:
437
+ pass
438
+
439
+ main()
stage2_object_v2/added_tokens.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c0284b582e14987fbd3d5a2cb2bd139084371ed9acbae488829a1c900833c680
3
+ size 707
stage2_object_v2/chat_template.jinja ADDED
@@ -0,0 +1,120 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {%- if tools %}
2
+ {{- '<|im_start|>system\n' }}
3
+ {%- if messages[0].role == 'system' %}
4
+ {%- if messages[0].content is string %}
5
+ {{- messages[0].content }}
6
+ {%- else %}
7
+ {%- for content in messages[0].content %}
8
+ {%- if 'text' in content %}
9
+ {{- content.text }}
10
+ {%- endif %}
11
+ {%- endfor %}
12
+ {%- endif %}
13
+ {{- '\n\n' }}
14
+ {%- endif %}
15
+ {{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
16
+ {%- for tool in tools %}
17
+ {{- "\n" }}
18
+ {{- tool | tojson }}
19
+ {%- endfor %}
20
+ {{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
21
+ {%- else %}
22
+ {%- if messages[0].role == 'system' %}
23
+ {{- '<|im_start|>system\n' }}
24
+ {%- if messages[0].content is string %}
25
+ {{- messages[0].content }}
26
+ {%- else %}
27
+ {%- for content in messages[0].content %}
28
+ {%- if 'text' in content %}
29
+ {{- content.text }}
30
+ {%- endif %}
31
+ {%- endfor %}
32
+ {%- endif %}
33
+ {{- '<|im_end|>\n' }}
34
+ {%- endif %}
35
+ {%- endif %}
36
+ {%- set image_count = namespace(value=0) %}
37
+ {%- set video_count = namespace(value=0) %}
38
+ {%- for message in messages %}
39
+ {%- if message.role == "user" %}
40
+ {{- '<|im_start|>' + message.role + '\n' }}
41
+ {%- if message.content is string %}
42
+ {{- message.content }}
43
+ {%- else %}
44
+ {%- for content in message.content %}
45
+ {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}
46
+ {%- set image_count.value = image_count.value + 1 %}
47
+ {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}
48
+ <|vision_start|><|image_pad|><|vision_end|>
49
+ {%- elif content.type == 'video' or 'video' in content %}
50
+ {%- set video_count.value = video_count.value + 1 %}
51
+ {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}
52
+ <|vision_start|><|video_pad|><|vision_end|>
53
+ {%- elif 'text' in content %}
54
+ {{- content.text }}
55
+ {%- endif %}
56
+ {%- endfor %}
57
+ {%- endif %}
58
+ {{- '<|im_end|>\n' }}
59
+ {%- elif message.role == "assistant" %}
60
+ {{- '<|im_start|>' + message.role + '\n' }}
61
+ {%- if message.content is string %}
62
+ {{- message.content }}
63
+ {%- else %}
64
+ {%- for content_item in message.content %}
65
+ {%- if 'text' in content_item %}
66
+ {{- content_item.text }}
67
+ {%- endif %}
68
+ {%- endfor %}
69
+ {%- endif %}
70
+ {%- if message.tool_calls %}
71
+ {%- for tool_call in message.tool_calls %}
72
+ {%- if (loop.first and message.content) or (not loop.first) %}
73
+ {{- '\n' }}
74
+ {%- endif %}
75
+ {%- if tool_call.function %}
76
+ {%- set tool_call = tool_call.function %}
77
+ {%- endif %}
78
+ {{- '<tool_call>\n{"name": "' }}
79
+ {{- tool_call.name }}
80
+ {{- '", "arguments": ' }}
81
+ {%- if tool_call.arguments is string %}
82
+ {{- tool_call.arguments }}
83
+ {%- else %}
84
+ {{- tool_call.arguments | tojson }}
85
+ {%- endif %}
86
+ {{- '}\n</tool_call>' }}
87
+ {%- endfor %}
88
+ {%- endif %}
89
+ {{- '<|im_end|>\n' }}
90
+ {%- elif message.role == "tool" %}
91
+ {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
92
+ {{- '<|im_start|>user' }}
93
+ {%- endif %}
94
+ {{- '\n<tool_response>\n' }}
95
+ {%- if message.content is string %}
96
+ {{- message.content }}
97
+ {%- else %}
98
+ {%- for content in message.content %}
99
+ {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}
100
+ {%- set image_count.value = image_count.value + 1 %}
101
+ {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}
102
+ <|vision_start|><|image_pad|><|vision_end|>
103
+ {%- elif content.type == 'video' or 'video' in content %}
104
+ {%- set video_count.value = video_count.value + 1 %}
105
+ {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}
106
+ <|vision_start|><|video_pad|><|vision_end|>
107
+ {%- elif 'text' in content %}
108
+ {{- content.text }}
109
+ {%- endif %}
110
+ {%- endfor %}
111
+ {%- endif %}
112
+ {{- '\n</tool_response>' }}
113
+ {%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
114
+ {{- '<|im_end|>\n' }}
115
+ {%- endif %}
116
+ {%- endif %}
117
+ {%- endfor %}
118
+ {%- if add_generation_prompt %}
119
+ {{- '<|im_start|>assistant\n' }}
120
+ {%- endif %}
stage2_object_v2/config.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b321be22f463f0af477f83cf7e63c566bb15645299208c6656a57ece7b9ffa87
3
+ size 1613
stage2_object_v2/generation_config.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3ff2e83a0510cccccc85c8c96c7df4207985c098b526d798bc2bc68e50bb1a41
3
+ size 199
stage2_object_v2/latest ADDED
@@ -0,0 +1 @@
 
 
1
+ global_step420
stage2_object_v2/merges.txt ADDED
The diff for this file is too large to render. See raw diff
 
stage2_object_v2/model.safetensors.index.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:dbdf49cc47bf9a028d0b1aa914b25401527572f03b092c8dfd9d428c9172783f
3
+ size 67791
stage2_object_v2/preprocessor_config.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:93585062a80db5e8ca038efc7726a3e6411d9db948472d81d63c6303993be8c5
3
+ size 782
stage2_object_v2/rng_state_0.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7272d7347df7a3eda0dcde7c67b3d4f9ff25e61e90c9673fc43693fe2e45be2f
3
+ size 15365
stage2_object_v2/rng_state_1.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0eb1f7b1acfd920c8cdef076eee48ccfec7d3ed4d8c5d83c2592fbbb4a4f9b38
3
+ size 15429
stage2_object_v2/rng_state_2.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5664fd5079c69dc8edcc429ded884f07d45c771f27a5a65fe1ad751348f9e1de
3
+ size 15429
stage2_object_v2/rng_state_3.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0a3f17c4e252e5b7e99f3573c3bca992070f9ef436dd578ee9e88c333a94cef0
3
+ size 15429
stage2_object_v2/scheduler.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e3433917e330a5e58b0185a3812a9e2c943fa2c157c04fecc317779eb366e4e6
3
+ size 1465
stage2_object_v2/special_tokens_map.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:76862e765266b85aa9459767e33cbaf13970f327a0e88d1c65846c2ddd3a1ecd
3
+ size 613
stage2_object_v2/tokenizer.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:aeb13307a71acd8fe81861d94ad54ab689df773318809eed3cbe794b4492dae4
3
+ size 11422654
stage2_object_v2/tokenizer_config.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cf43a5bf1a49ee69ecced02f419b169e72559034dcf15af47cf775bd253830f0
3
+ size 5472
stage2_object_v2/trainer_state.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b0e46e697ee453aa2a87361fbd2616515614af85b224376d525af005dea2e8eb
3
+ size 9250
stage2_object_v2/training_args.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:314b5adbe54f49c9e00e177abc56e7953500d7fd77e68390e5a0608c2fb34a90
3
+ size 8209
stage2_object_v2/video_preprocessor_config.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:59c5c9eb52182eb14c06ffb10ca9effd29adce5f238a95de23ca14a38dbd2cb1
3
+ size 817
stage2_object_v2/vocab.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ca10d7e9fb3ed18575dd1e277a2579c16d108e32f27439684afa0e10b1440910
3
+ size 2776833
stage2_object_v2/zero_to_fp32.py ADDED
@@ -0,0 +1,760 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+
3
+ # Copyright (c) Microsoft Corporation.
4
+ # SPDX-License-Identifier: Apache-2.0
5
+
6
+ # DeepSpeed Team
7
+
8
+ # This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets
9
+ # copied into the top level checkpoint dir, so the user can easily do the conversion at any point in
10
+ # the future. Once extracted, the weights don't require DeepSpeed and can be used in any
11
+ # application.
12
+ #
13
+ # example:
14
+ # python zero_to_fp32.py . output_dir/
15
+ # or
16
+ # python zero_to_fp32.py . output_dir/ --safe_serialization
17
+
18
+ import argparse
19
+ import torch
20
+ import glob
21
+ import math
22
+ import os
23
+ import re
24
+ import gc
25
+ import json
26
+ import numpy as np
27
+ from tqdm import tqdm
28
+ from collections import OrderedDict
29
+ from dataclasses import dataclass
30
+
31
+ # while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with
32
+ # DeepSpeed data structures it has to be available in the current python environment.
33
+ from deepspeed.utils import logger
34
+ from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS,
35
+ FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES,
36
+ FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS)
37
+
38
+
39
+ @dataclass
40
+ class zero_model_state:
41
+ buffers: dict()
42
+ param_shapes: dict()
43
+ shared_params: list
44
+ ds_version: int
45
+ frozen_param_shapes: dict()
46
+ frozen_param_fragments: dict()
47
+
48
+
49
+ debug = 0
50
+
51
+ # load to cpu
52
+ device = torch.device('cpu')
53
+
54
+
55
+ def atoi(text):
56
+ return int(text) if text.isdigit() else text
57
+
58
+
59
+ def natural_keys(text):
60
+ '''
61
+ alist.sort(key=natural_keys) sorts in human order
62
+ http://nedbatchelder.com/blog/200712/human_sorting.html
63
+ (See Toothy's implementation in the comments)
64
+ '''
65
+ return [atoi(c) for c in re.split(r'(\d+)', text)]
66
+
67
+
68
+ def get_model_state_file(checkpoint_dir, zero_stage):
69
+ if not os.path.isdir(checkpoint_dir):
70
+ raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist")
71
+
72
+ # there should be only one file
73
+ if zero_stage <= 2:
74
+ file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt")
75
+ elif zero_stage == 3:
76
+ file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt")
77
+
78
+ if not os.path.exists(file):
79
+ raise FileNotFoundError(f"can't find model states file at '{file}'")
80
+
81
+ return file
82
+
83
+
84
+ def get_checkpoint_files(checkpoint_dir, glob_pattern):
85
+ # XXX: need to test that this simple glob rule works for multi-node setup too
86
+ ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys)
87
+
88
+ if len(ckpt_files) == 0:
89
+ raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'")
90
+
91
+ return ckpt_files
92
+
93
+
94
+ def get_optim_files(checkpoint_dir):
95
+ return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt")
96
+
97
+
98
+ def get_model_state_files(checkpoint_dir):
99
+ return get_checkpoint_files(checkpoint_dir, "*_model_states.pt")
100
+
101
+
102
+ def parse_model_states(files):
103
+ zero_model_states = []
104
+ for file in files:
105
+ state_dict = torch.load(file, map_location=device, weights_only=False)
106
+
107
+ if BUFFER_NAMES not in state_dict:
108
+ raise ValueError(f"{file} is not a model state checkpoint")
109
+ buffer_names = state_dict[BUFFER_NAMES]
110
+ if debug:
111
+ print("Found buffers:", buffer_names)
112
+
113
+ # recover just the buffers while restoring them to fp32 if they were saved in fp16
114
+ buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names}
115
+ param_shapes = state_dict[PARAM_SHAPES]
116
+
117
+ # collect parameters that are included in param_shapes
118
+ param_names = []
119
+ for s in param_shapes:
120
+ for name in s.keys():
121
+ param_names.append(name)
122
+
123
+ # update with frozen parameters
124
+ frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None)
125
+ if frozen_param_shapes is not None:
126
+ if debug:
127
+ print(f"Found frozen_param_shapes: {frozen_param_shapes}")
128
+ param_names += list(frozen_param_shapes.keys())
129
+
130
+ # handle shared params
131
+ shared_params = [[k, v] for k, v in state_dict["shared_params"].items()]
132
+
133
+ ds_version = state_dict.get(DS_VERSION, None)
134
+
135
+ frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None)
136
+
137
+ z_model_state = zero_model_state(buffers=buffers,
138
+ param_shapes=param_shapes,
139
+ shared_params=shared_params,
140
+ ds_version=ds_version,
141
+ frozen_param_shapes=frozen_param_shapes,
142
+ frozen_param_fragments=frozen_param_fragments)
143
+ zero_model_states.append(z_model_state)
144
+
145
+ return zero_model_states
146
+
147
+
148
+ def parse_optim_states(files, ds_checkpoint_dir):
149
+ total_files = len(files)
150
+ state_dicts = []
151
+ for f in tqdm(files, desc='Loading checkpoint shards'):
152
+ state_dict = torch.load(f, map_location=device, mmap=True, weights_only=False)
153
+ # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights
154
+ # and also handle the case where it was already removed by another helper script
155
+ state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None)
156
+ state_dicts.append(state_dict)
157
+
158
+ if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]:
159
+ raise ValueError(f"{files[0]} is not a zero checkpoint")
160
+ zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE]
161
+ world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT]
162
+
163
+ # For ZeRO-2 each param group can have different partition_count as data parallelism for expert
164
+ # parameters can be different from data parallelism for non-expert parameters. So we can just
165
+ # use the max of the partition_count to get the dp world_size.
166
+
167
+ if type(world_size) is list:
168
+ world_size = max(world_size)
169
+
170
+ if world_size != total_files:
171
+ raise ValueError(
172
+ f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. "
173
+ "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes."
174
+ )
175
+
176
+ # the groups are named differently in each stage
177
+ if zero_stage <= 2:
178
+ fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS
179
+ elif zero_stage == 3:
180
+ fp32_groups_key = FP32_FLAT_GROUPS
181
+ else:
182
+ raise ValueError(f"unknown zero stage {zero_stage}")
183
+
184
+ fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))]
185
+ return zero_stage, world_size, fp32_flat_groups
186
+
187
+
188
+ def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters):
189
+ """
190
+ Returns fp32 state_dict reconstructed from ds checkpoint
191
+
192
+ Args:
193
+ - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are)
194
+
195
+ """
196
+ print(f"Processing zero checkpoint '{ds_checkpoint_dir}'")
197
+
198
+ optim_files = get_optim_files(ds_checkpoint_dir)
199
+ zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir)
200
+ print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}")
201
+
202
+ model_files = get_model_state_files(ds_checkpoint_dir)
203
+
204
+ zero_model_states = parse_model_states(model_files)
205
+ print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}')
206
+
207
+ if zero_stage <= 2:
208
+ return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states,
209
+ exclude_frozen_parameters)
210
+ elif zero_stage == 3:
211
+ return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states,
212
+ exclude_frozen_parameters)
213
+
214
+
215
+ def _zero2_merge_frozen_params(state_dict, zero_model_states):
216
+ if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0:
217
+ return
218
+
219
+ frozen_param_shapes = zero_model_states[0].frozen_param_shapes
220
+ frozen_param_fragments = zero_model_states[0].frozen_param_fragments
221
+
222
+ if debug:
223
+ num_elem = sum(s.numel() for s in frozen_param_shapes.values())
224
+ print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}')
225
+
226
+ wanted_params = len(frozen_param_shapes)
227
+ wanted_numel = sum(s.numel() for s in frozen_param_shapes.values())
228
+ avail_numel = sum([p.numel() for p in frozen_param_fragments.values()])
229
+ print(f'Frozen params: Have {avail_numel} numels to process.')
230
+ print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params')
231
+
232
+ total_params = 0
233
+ total_numel = 0
234
+ for name, shape in frozen_param_shapes.items():
235
+ total_params += 1
236
+ unpartitioned_numel = shape.numel()
237
+ total_numel += unpartitioned_numel
238
+
239
+ state_dict[name] = frozen_param_fragments[name]
240
+
241
+ if debug:
242
+ print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ")
243
+
244
+ print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements")
245
+
246
+
247
+ def _has_callable(obj, fn):
248
+ attr = getattr(obj, fn, None)
249
+ return callable(attr)
250
+
251
+
252
+ def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states):
253
+ param_shapes = zero_model_states[0].param_shapes
254
+
255
+ # Reconstruction protocol:
256
+ #
257
+ # XXX: document this
258
+
259
+ if debug:
260
+ for i in range(world_size):
261
+ for j in range(len(fp32_flat_groups[0])):
262
+ print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}")
263
+
264
+ # XXX: memory usage doubles here (zero2)
265
+ num_param_groups = len(fp32_flat_groups[0])
266
+ merged_single_partition_of_fp32_groups = []
267
+ for i in range(num_param_groups):
268
+ merged_partitions = [sd[i] for sd in fp32_flat_groups]
269
+ full_single_fp32_vector = torch.cat(merged_partitions, 0)
270
+ merged_single_partition_of_fp32_groups.append(full_single_fp32_vector)
271
+ avail_numel = sum(
272
+ [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups])
273
+
274
+ if debug:
275
+ wanted_params = sum([len(shapes) for shapes in param_shapes])
276
+ wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes])
277
+ # not asserting if there is a mismatch due to possible padding
278
+ print(f"Have {avail_numel} numels to process.")
279
+ print(f"Need {wanted_numel} numels in {wanted_params} params.")
280
+
281
+ # params
282
+ # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support
283
+ # out-of-core computing solution
284
+ total_numel = 0
285
+ total_params = 0
286
+ for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups):
287
+ offset = 0
288
+ avail_numel = full_single_fp32_vector.numel()
289
+ for name, shape in shapes.items():
290
+
291
+ unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape)
292
+ total_numel += unpartitioned_numel
293
+ total_params += 1
294
+
295
+ if debug:
296
+ print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ")
297
+ state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape)
298
+ offset += unpartitioned_numel
299
+
300
+ # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and
301
+ # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex
302
+ # paddings performed in the code it's almost impossible to predict the exact numbers w/o the
303
+ # live optimizer object, so we are checking that the numbers are within the right range
304
+ align_to = 2 * world_size
305
+
306
+ def zero2_align(x):
307
+ return align_to * math.ceil(x / align_to)
308
+
309
+ if debug:
310
+ print(f"original offset={offset}, avail_numel={avail_numel}")
311
+
312
+ offset = zero2_align(offset)
313
+ avail_numel = zero2_align(avail_numel)
314
+
315
+ if debug:
316
+ print(f"aligned offset={offset}, avail_numel={avail_numel}")
317
+
318
+ # Sanity check
319
+ if offset != avail_numel:
320
+ raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong")
321
+
322
+ print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements")
323
+
324
+
325
+ def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states,
326
+ exclude_frozen_parameters):
327
+ state_dict = OrderedDict()
328
+
329
+ # buffers
330
+ buffers = zero_model_states[0].buffers
331
+ state_dict.update(buffers)
332
+ if debug:
333
+ print(f"added {len(buffers)} buffers")
334
+
335
+ if not exclude_frozen_parameters:
336
+ _zero2_merge_frozen_params(state_dict, zero_model_states)
337
+
338
+ _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states)
339
+
340
+ # recover shared parameters
341
+ for pair in zero_model_states[0].shared_params:
342
+ if pair[1] in state_dict:
343
+ state_dict[pair[0]] = state_dict[pair[1]]
344
+
345
+ return state_dict
346
+
347
+
348
+ def zero3_partitioned_param_info(unpartitioned_numel, world_size):
349
+ remainder = unpartitioned_numel % world_size
350
+ padding_numel = (world_size - remainder) if remainder else 0
351
+ partitioned_numel = math.ceil(unpartitioned_numel / world_size)
352
+ return partitioned_numel, padding_numel
353
+
354
+
355
+ def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states):
356
+ if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0:
357
+ return
358
+
359
+ if debug:
360
+ for i in range(world_size):
361
+ num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values())
362
+ print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}')
363
+
364
+ frozen_param_shapes = zero_model_states[0].frozen_param_shapes
365
+ wanted_params = len(frozen_param_shapes)
366
+ wanted_numel = sum(s.numel() for s in frozen_param_shapes.values())
367
+ avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size
368
+ print(f'Frozen params: Have {avail_numel} numels to process.')
369
+ print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params')
370
+
371
+ total_params = 0
372
+ total_numel = 0
373
+ for name, shape in zero_model_states[0].frozen_param_shapes.items():
374
+ total_params += 1
375
+ unpartitioned_numel = shape.numel()
376
+ total_numel += unpartitioned_numel
377
+
378
+ param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states)
379
+ state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape)
380
+
381
+ partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size)
382
+
383
+ if debug:
384
+ print(
385
+ f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}"
386
+ )
387
+
388
+ print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements")
389
+
390
+
391
+ class GatheredTensor:
392
+ """
393
+ A pseudo tensor that collects partitioned weights.
394
+ It is more memory efficient when there are multiple groups.
395
+ """
396
+
397
+ def __init__(self, flat_groups, flat_groups_offset, offset, partitioned_numel, shape):
398
+ self.flat_groups = flat_groups
399
+ self.flat_groups_offset = flat_groups_offset
400
+ self.offset = offset
401
+ self.partitioned_numel = partitioned_numel
402
+ self.shape = shape
403
+ self.dtype = self.flat_groups[0][0].dtype
404
+
405
+ def contiguous(self):
406
+ """
407
+ Merge partitioned weights from flat_groups into a single tensor.
408
+ """
409
+ end_idx = self.offset + self.partitioned_numel
410
+ world_size = len(self.flat_groups)
411
+ pad_flat_param_chunks = []
412
+
413
+ for rank_i in range(world_size):
414
+ # for each rank, we need to collect weights from related group/groups
415
+ flat_groups_at_rank_i = self.flat_groups[rank_i]
416
+ start_group_id = None
417
+ end_group_id = None
418
+ for group_id in range(len(self.flat_groups_offset)):
419
+ if self.flat_groups_offset[group_id] <= self.offset < self.flat_groups_offset[group_id + 1]:
420
+ start_group_id = group_id
421
+ if self.flat_groups_offset[group_id] < end_idx <= self.flat_groups_offset[group_id + 1]:
422
+ end_group_id = group_id
423
+ break
424
+ # collect weights from related group/groups
425
+ for group_id in range(start_group_id, end_group_id + 1):
426
+ flat_tensor = flat_groups_at_rank_i[group_id]
427
+ start_offset = self.offset - self.flat_groups_offset[group_id]
428
+ end_offset = min(end_idx, self.flat_groups_offset[group_id + 1]) - self.flat_groups_offset[group_id]
429
+ pad_flat_param_chunks.append(flat_tensor[start_offset:end_offset])
430
+
431
+ # collect weights from all ranks
432
+ pad_flat_param = torch.cat(pad_flat_param_chunks, dim=0)
433
+ param = pad_flat_param[:self.shape.numel()].view(self.shape).contiguous()
434
+ return param
435
+
436
+
437
+ def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states):
438
+ param_shapes = zero_model_states[0].param_shapes
439
+ avail_numel = sum([flat_group.numel() for flat_group in fp32_flat_groups[0]]) * world_size
440
+
441
+ # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each
442
+ # param, re-consolidating each param, while dealing with padding if any
443
+
444
+ # merge list of dicts, preserving order
445
+ param_shapes = {k: v for d in param_shapes for k, v in d.items()}
446
+
447
+ if debug:
448
+ for i in range(world_size):
449
+ print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}")
450
+
451
+ wanted_params = len(param_shapes)
452
+ wanted_numel = sum(shape.numel() for shape in param_shapes.values())
453
+ # not asserting if there is a mismatch due to possible padding
454
+ avail_numel = fp32_flat_groups[0].numel() * world_size
455
+ print(f"Trainable params: Have {avail_numel} numels to process.")
456
+ print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.")
457
+
458
+ # params
459
+ # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support
460
+ # out-of-core computing solution
461
+ offset = 0
462
+ total_numel = 0
463
+ total_params = 0
464
+ flat_groups_offset = [0] + list(np.cumsum([flat_tensor.numel() for flat_tensor in fp32_flat_groups[0]]))
465
+ for name, shape in tqdm(param_shapes.items(), desc='Gathering sharded weights'):
466
+ unpartitioned_numel = shape.numel()
467
+ total_numel += unpartitioned_numel
468
+ total_params += 1
469
+ partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size)
470
+
471
+ if debug:
472
+ print(
473
+ f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}"
474
+ )
475
+
476
+ # memory efficient tensor
477
+ tensor = GatheredTensor(fp32_flat_groups, flat_groups_offset, offset, partitioned_numel, shape)
478
+ state_dict[name] = tensor
479
+ offset += partitioned_numel
480
+
481
+ offset *= world_size
482
+
483
+ # Sanity check
484
+ if offset != avail_numel:
485
+ raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong")
486
+
487
+ print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements")
488
+
489
+
490
+ def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states,
491
+ exclude_frozen_parameters):
492
+ state_dict = OrderedDict()
493
+
494
+ # buffers
495
+ buffers = zero_model_states[0].buffers
496
+ state_dict.update(buffers)
497
+ if debug:
498
+ print(f"added {len(buffers)} buffers")
499
+
500
+ if not exclude_frozen_parameters:
501
+ _zero3_merge_frozen_params(state_dict, world_size, zero_model_states)
502
+
503
+ _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states)
504
+
505
+ # recover shared parameters
506
+ for pair in zero_model_states[0].shared_params:
507
+ if pair[1] in state_dict:
508
+ state_dict[pair[0]] = state_dict[pair[1]]
509
+
510
+ return state_dict
511
+
512
+
513
+ def to_torch_tensor(state_dict, return_empty_tensor=False):
514
+ """
515
+ Convert state_dict of GatheredTensor to torch tensor
516
+ """
517
+ torch_state_dict = {}
518
+ converted_tensors = {}
519
+ for name, tensor in state_dict.items():
520
+ tensor_id = id(tensor)
521
+ if tensor_id in converted_tensors: # shared tensors
522
+ shared_tensor = torch_state_dict[converted_tensors[tensor_id]]
523
+ torch_state_dict[name] = shared_tensor
524
+ else:
525
+ converted_tensors[tensor_id] = name
526
+ if return_empty_tensor:
527
+ torch_state_dict[name] = torch.empty(tensor.shape, dtype=tensor.dtype)
528
+ else:
529
+ torch_state_dict[name] = tensor.contiguous()
530
+ return torch_state_dict
531
+
532
+
533
+ def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir,
534
+ tag=None,
535
+ exclude_frozen_parameters=False,
536
+ lazy_mode=False):
537
+ """
538
+ Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with
539
+ ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example
540
+ via a model hub.
541
+
542
+ Args:
543
+ - ``checkpoint_dir``: path to the desired checkpoint folder
544
+ - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14``
545
+ - ``exclude_frozen_parameters``: exclude frozen parameters
546
+ - ``lazy_mode``: get state_dict in lazy mode. It returns a dict of pesduo tensor instead of torch tensor, which is more memory efficient.
547
+ Convert the pesduo tensor to torch tensor by ``.contiguous()``
548
+
549
+ Returns:
550
+ - pytorch ``state_dict``
551
+
552
+ A typical usage might be ::
553
+
554
+ from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint
555
+ # do the training and checkpoint saving
556
+ state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu
557
+ model = model.cpu() # move to cpu
558
+ model.load_state_dict(state_dict)
559
+ # submit to model hub or save the model to share with others
560
+
561
+ In this example the ``model`` will no longer be usable in the deepspeed context of the same
562
+ application. i.e. you will need to re-initialize the deepspeed engine, since
563
+ ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it.
564
+
565
+ If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead.
566
+
567
+ Note: the above usage may not work if your application doesn't have sufficient free CPU memory.
568
+ You may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with
569
+ the checkpoint. Or you can load state_dict in lazy mode ::
570
+
571
+ from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint
572
+ state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, lazy_mode=True) # not on cpu
573
+ for name, lazy_tensor in state_dict.item():
574
+ tensor = lazy_tensor.contiguous() # to cpu
575
+ print(name, tensor)
576
+ # del tensor to release memory if it no longer in use
577
+ """
578
+ if tag is None:
579
+ latest_path = os.path.join(checkpoint_dir, 'latest')
580
+ if os.path.isfile(latest_path):
581
+ with open(latest_path, 'r') as fd:
582
+ tag = fd.read().strip()
583
+ else:
584
+ raise ValueError(f"Unable to find 'latest' file at {latest_path}")
585
+
586
+ ds_checkpoint_dir = os.path.join(checkpoint_dir, tag)
587
+
588
+ if not os.path.isdir(ds_checkpoint_dir):
589
+ raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist")
590
+
591
+ state_dict = _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters)
592
+ if lazy_mode:
593
+ return state_dict
594
+ else:
595
+ return to_torch_tensor(state_dict)
596
+
597
+
598
+ def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir,
599
+ output_dir,
600
+ max_shard_size="5GB",
601
+ safe_serialization=False,
602
+ tag=None,
603
+ exclude_frozen_parameters=False):
604
+ """
605
+ Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be
606
+ loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed.
607
+
608
+ Args:
609
+ - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``)
610
+ - ``output_dir``: directory to the pytorch fp32 state_dict output files
611
+ - ``max_shard_size``: the maximum size for a checkpoint before being sharded, default value is 5GB
612
+ - ``safe_serialization``: whether to save the model using `safetensors` or the traditional PyTorch way (that uses `pickle`).
613
+ - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14``
614
+ - ``exclude_frozen_parameters``: exclude frozen parameters
615
+ """
616
+
617
+ # Dependency pre-check
618
+ if safe_serialization:
619
+ try:
620
+ from safetensors.torch import save_file
621
+ except ImportError:
622
+ print('If you want to use `safe_serialization`, please `pip install safetensors`')
623
+ raise
624
+ if max_shard_size is not None:
625
+ try:
626
+ from huggingface_hub import split_torch_state_dict_into_shards
627
+ except ImportError:
628
+ print('If you want to use `max_shard_size`, please `pip install huggingface_hub`')
629
+ raise
630
+
631
+ # Convert zero checkpoint to state_dict
632
+ state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir,
633
+ tag,
634
+ exclude_frozen_parameters,
635
+ lazy_mode=True)
636
+
637
+ # Shard the model if it is too big.
638
+ weights_name = "model.safetensors" if safe_serialization else "pytorch_model.bin"
639
+ if max_shard_size is not None:
640
+ filename_pattern = weights_name.replace(".bin", "{suffix}.bin").replace(".safetensors", "{suffix}.safetensors")
641
+ # an memory-efficient approach for sharding
642
+ empty_state_dict = to_torch_tensor(state_dict, return_empty_tensor=True)
643
+ state_dict_split = split_torch_state_dict_into_shards(empty_state_dict,
644
+ filename_pattern=filename_pattern,
645
+ max_shard_size=max_shard_size)
646
+ else:
647
+ from collections import namedtuple
648
+ StateDictSplit = namedtuple("StateDictSplit", ["is_sharded", "filename_to_tensors"])
649
+ state_dict_split = StateDictSplit(is_sharded=False,
650
+ filename_to_tensors={weights_name: list(state_dict.keys())})
651
+
652
+ # Save the model by shard
653
+ os.makedirs(output_dir, exist_ok=True)
654
+ filename_to_tensors = state_dict_split.filename_to_tensors.items()
655
+ for shard_file, tensors in tqdm(filename_to_tensors, desc="Saving checkpoint shards"):
656
+ shard_state_dict = {tensor_name: state_dict[tensor_name] for tensor_name in tensors}
657
+ shard_state_dict = to_torch_tensor(shard_state_dict)
658
+ output_path = os.path.join(output_dir, shard_file)
659
+ if safe_serialization:
660
+ save_file(shard_state_dict, output_path, metadata={"format": "pt"})
661
+ else:
662
+ torch.save(shard_state_dict, output_path)
663
+ # release the memory of current shard
664
+ for tensor_name in list(shard_state_dict.keys()):
665
+ del state_dict[tensor_name]
666
+ del shard_state_dict[tensor_name]
667
+ del shard_state_dict
668
+ gc.collect()
669
+
670
+ # Save index if sharded
671
+ if state_dict_split.is_sharded:
672
+ index = {
673
+ "metadata": state_dict_split.metadata,
674
+ "weight_map": state_dict_split.tensor_to_filename,
675
+ }
676
+ save_index_file = "model.safetensors.index.json" if safe_serialization else "pytorch_model.bin.index.json"
677
+ save_index_file = os.path.join(output_dir, save_index_file)
678
+ with open(save_index_file, "w", encoding="utf-8") as f:
679
+ content = json.dumps(index, indent=2, sort_keys=True) + "\n"
680
+ f.write(content)
681
+
682
+
683
+ def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None):
684
+ """
685
+ 1. Put the provided model to cpu
686
+ 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict``
687
+ 3. Load it into the provided model
688
+
689
+ Args:
690
+ - ``model``: the model object to update
691
+ - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``)
692
+ - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14``
693
+
694
+ Returns:
695
+ - ``model`: modified model
696
+
697
+ Make sure you have plenty of CPU memory available before you call this function. If you don't
698
+ have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it
699
+ conveniently placed for you in the checkpoint folder.
700
+
701
+ A typical usage might be ::
702
+
703
+ from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint
704
+ model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir)
705
+ # submit to model hub or save the model to share with others
706
+
707
+ Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context
708
+ of the same application. i.e. you will need to re-initialize the deepspeed engine, since
709
+ ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it.
710
+
711
+ """
712
+ logger.info(f"Extracting fp32 weights")
713
+ state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag)
714
+
715
+ logger.info(f"Overwriting model with fp32 weights")
716
+ model = model.cpu()
717
+ model.load_state_dict(state_dict, strict=False)
718
+
719
+ return model
720
+
721
+
722
+ if __name__ == "__main__":
723
+ parser = argparse.ArgumentParser()
724
+ parser.add_argument("checkpoint_dir",
725
+ type=str,
726
+ help="path to the desired checkpoint folder, e.g., path/checkpoint-12")
727
+ parser.add_argument("output_dir",
728
+ type=str,
729
+ help="directory to the pytorch fp32 state_dict output files"
730
+ "(e.g. path/checkpoint-12-output/)")
731
+ parser.add_argument(
732
+ "--max_shard_size",
733
+ type=str,
734
+ default="5GB",
735
+ help="The maximum size for a checkpoint before being sharded. Checkpoints shard will then be each of size"
736
+ "lower than this size. If expressed as a string, needs to be digits followed by a unit (like `5MB`"
737
+ "We default it to 5GB in order for models to be able to run easily on free-tier google colab instances"
738
+ "without CPU OOM issues.")
739
+ parser.add_argument(
740
+ "--safe_serialization",
741
+ default=False,
742
+ action='store_true',
743
+ help="Whether to save the model using `safetensors` or the traditional PyTorch way (that uses `pickle`).")
744
+ parser.add_argument("-t",
745
+ "--tag",
746
+ type=str,
747
+ default=None,
748
+ help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1")
749
+ parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters")
750
+ parser.add_argument("-d", "--debug", action='store_true', help="enable debug")
751
+ args = parser.parse_args()
752
+
753
+ debug = args.debug
754
+
755
+ convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir,
756
+ args.output_dir,
757
+ max_shard_size=args.max_shard_size,
758
+ safe_serialization=args.safe_serialization,
759
+ tag=args.tag,
760
+ exclude_frozen_parameters=args.exclude_frozen_parameters)