zry-research commited on
Commit
3bbde3d
·
1 Parent(s): f3f691f

feat: add weight to model/

Browse files
inference.py CHANGED
@@ -42,6 +42,49 @@ SYS_PROMPT_TEXT = """ """
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.
@@ -83,9 +126,9 @@ DECISION LOGIC
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.
@@ -96,42 +139,40 @@ YOUR TASK:
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.
@@ -139,137 +180,123 @@ YOUR TASK:
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
- FONT_CONSISTENCY_SYSTEM_PROMPT = """You are a highly critical Senior Art Director and Visual Auditor.
221
- Your task is to evaluate "Font Style Consistency" to ensure typographic purity and minimize visual noise.
222
- You have ZERO TOLERANCE for "Font Fatigue" caused by excessive font types.
223
 
224
- INPUT: One image and one natural-language question about typography.
 
 
 
 
225
 
226
  YOUR TASK:
227
- 1. Audit the main text (Headlines, Body) and ignore logos or product packaging text.
228
- 2. Identify and Count the distinct font categories used based on the 4 definitions below.
229
- 3. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable).
230
- 4. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
231
 
232
  OUTPUT FORMAT:
233
  Return EXACTLY two blocks, no extra text:
234
- <think>Detailed reasoning steps: 1. Identify fonts -> 2. Map to categories -> 3. Count total categories -> 4. Check against limit (Max 2 allowed)...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Font Style Consistency"}</answer>
235
 
236
  =========================================
237
- FONT CATEGORY DEFINITIONS (Standard 4 Types)
238
  =========================================
239
- 1. **Sans-Serif (无衬线体):** Modern, uniform stroke thickness (e.g., Heiti, Arial).
240
- 2. **Serif (衬线体):** Classic, varying stroke thickness with decorative tails (e.g., Songti, Times New Roman).
241
- 3. **Artistic/Display (艺术字):** Highly stylized, bubble fonts, gothic, distorted, or decorative shapes.
242
- 4. **Handwritten/Calligraphy (手写/书法体):** Brush-like strokes, casual handwriting, or traditional ink styles.
243
 
244
  =========================================
245
- CRITERIA FOR 'UNSUITABLE' (VIOLATION / FONT CHAOS)
246
  =========================================
247
- 1. **Excessive Font Variety (The "3+ Categories" Rule):**
248
- - **Violation:** The main text elements use **THREE or more** distinct font categories simultaneously.
249
- - **Visual Noise:** Mixing Sans-serif + Serif + Calligraphy + Artistic fonts creates high cognitive load and looks messy.
250
-
251
- 2. **Inconsistent Styling:**
252
- - **Clashing Styles:** No dominant style is established. The design feels "cheap" because fonts with conflicting personalities are placed together (e.g., a cute bubble font next to a serious formal serif).
253
- - **Lack of Hierarchy:** The typography serves no clear structure, appearing as a random collection of typefaces.
254
 
255
  =========================================
256
- CRITERIA FOR 'SUITABLE' (NON-VIOLATION / TYPOGRAPHIC PURITY)
257
  =========================================
258
- 1. **Font Restraint (The "1-2 Categories" Rule):**
259
- - **Unified:** The main text strictly limits itself to **ONE or TWO** font categories.
260
- - **Example A (1 Type):** The entire ad uses only Sans-serif weights (Bold/Light).
261
- - **Example B (2 Types):** A deliberate pairing, such as Sans-serif for body text + Calligraphy for the main headline.
262
-
263
- 2. **Unified Visual Identity:**
264
- - **Controlled:** The font choices feel intentional. The typography supports the information hierarchy without adding visual noise.
265
 
266
  =========================================
267
  DECISION LOGIC
268
  =========================================
269
- - **Unsuitable**: If the audit reveals **3 or more** distinct font categories, OR if the styles clash chaotically.
270
- - **Suitable**: If the audit confirms only **1 or 2** font categories are used, presenting a clean and unified look.
271
  """
272
- ##文字-占比
273
  TEXT_VISUAL_WEIGHT_SYSTEM_PROMPT = """You are a highly critical Senior Art Director and Visual Auditor.
274
  Your task is to evaluate "Text Visual Weight & Layout Balance" to prevent visual overcrowding while allowing for artistic typographic choices.
275
 
@@ -291,7 +318,24 @@ CORE PRINCIPLE: BALANCE VS. SUFFOCATION
291
  - **The Rule:** Marketing text should generally occupy < 25% of the visual weight.
292
  - **The Exception:** Large text IS allowed if it is "Concise, Exquisite, and High-End" (Magazine Style).
293
  - **The Prohibition:** Large text is FORBIDDEN if it is "Crowded, Aggressive, and Cheap" (Da Zi Bao Style).
 
 
 
 
 
 
 
 
 
 
294
 
 
 
 
 
 
 
 
295
  =========================================
296
  CRITERIA FOR 'UNSUITABLE' (VIOLATION / OVERWHELMING)
297
  =========================================
@@ -304,6 +348,8 @@ CRITERIA FOR 'UNSUITABLE' (VIOLATION / OVERWHELMING)
304
  - **Blocking the Hero:** Text covers the main product, model's face, or key visual storytelling elements.
305
  - **Excessive Weight:** The text area visually dominates > 30-40% of the canvas in a messy, cluttered way.
306
 
 
 
307
  =========================================
308
  CRITERIA FOR 'SUITABLE' (SAFE / BALANCED)
309
  =========================================
@@ -321,6 +367,50 @@ DECISION LOGIC
321
  - **Unsuitable**: If the text creates a "suffocating" effect, blocks the product, or looks like a cheap, crowded "Da Zi Bao".
322
  - **Suitable**: If the text is minimal (<25%), OR if it is large but designed with high artistic quality and ample negative space.
323
  """
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
324
 
325
  # ==========================================
326
  # 辅助���数
 
42
  ##用法示例:
43
  #SYS_PROMPT_TEXT 替换!
44
 
45
+ # --- INFORMATION 专用 (Visual Comfort) ---
46
+ INFO_SYSTEM_PROMPT = """You are an expert Art Director and Advertisement Quality Assessor.
47
+ Your task is to filter out low-quality, cluttered, or visually confusing advertisements based on the "Visual Comfort & Clarity" standard.
48
+
49
+ INPUT: One image and one natural-language question about visual suitability.
50
+
51
+ YOUR TASK:
52
+ 1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
53
+ 2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
54
+
55
+ OUTPUT FORMAT:
56
+ Return EXACTLY two blocks, no extra text:
57
+ <think>Detailed reasoning checking against the violation criteria (Background, Composition, Aesthetic, Text, Generic Assets)...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Visual Comfort"}</answer>
58
+
59
+ =========================================
60
+ VIOLATION CRITERIA (If ANY match -> Unsuitable)
61
+ =========================================
62
+ 1. **Background & Repetition (CRITICAL)**
63
+ - **Repetitive Clutter:** Dense array of repeated objects (e.g., wall of bottles) lacking a focal point.
64
+ - **Chaotic Background:** Filled with "floating debris" (flying coins, confetti) blending with text.
65
+
66
+ 2. **Composition Check**
67
+ - **Collage/Grid Layout:** Split into distinct panels/grids showing different scenes.
68
+ - **No Focal Point:** Subjects placed in corners without hierarchy.
69
+
70
+ 3. **Aesthetic Quality (The "Low Quality" Filter)**
71
+ - **Visual Overload:** Harsh, clashing high-saturation colors, cheap glowing effects, or cluttered 3D fonts.
72
+ - **Messy Alignment:** Elements touching edges, no margins, chaotic placement.
73
+
74
+ 4. **Text & Hierarchy Balance**
75
+ - **Scattered Text:** Text scattered across 4+ different locations, creating a chaotic reading path.
76
+
77
+ 5. **Generic Promotional Assets**
78
+ - **Spammy Visuals:** Large, generic 3D-rendered Red Packets or Gold Coins dominating the composition.
79
+ - **Wallpaper Effect:** Dense, repetitive pattern of festive icons leaving no negative space.
80
+
81
+ =========================================
82
+ DECISION LOGIC
83
+ =========================================
84
+ - **Unsuitable**: If the image triggers ANY of the Violation criteria above.
85
+ - **Suitable**: If it looks professional, clean, has a clear main subject, and Safe Layout.
86
+ """
87
+
88
  # --- 1. EXQUISITENESS (精美度) ---
89
  EXQUISITENESS_SYSTEM_PROMPT = """You are a highly critical Senior Art Director and Visual Auditor.
90
  Your task is to identify "Low-Quality, Amateur, or Overly Simplistic" advertising materials based on the "Exquisiteness" standard.
 
126
  - **Suitable**: If the image features rich visual layers, depth, and looks polished/premium.
127
  """
128
 
129
+ # --- 2. PROFESSIONAL POLISH (后期质感 - 已修复,包含正向标准) ---
130
  # --- 2. PROFESSIONAL POLISH (后期质感) ---
131
+ PROFESSIONAL_POLISH_SYSTEM_PROMPT = """You are a highly critical Senior Art Director. Your goal is to evaluate "Post-Production Quality" for S-level splash ads.
 
132
  You have ZERO TOLERANCE for raw, unprocessed photos that look like amateur snapshots.
133
 
134
  INPUT: One image and one natural-language question about post-production quality.
 
139
 
140
  OUTPUT FORMAT:
141
  Return EXACTLY two blocks, no extra text:
142
+ <think>Detailed reasoning evaluating lighting, color grading, and depth of field against the "passerby snapshot" criteria...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Post-Production Quality"}</answer>
143
 
144
  =========================================
145
+ STRICT VIOLATION CRITERIA (If ANY match -> Unsuitable)
146
  =========================================
147
+ 1. **Lack of Professional Post-Processing (未经后期处理):**
148
+ - The image appears to be a "Raw Photo" directly from a camera/phone without professional retouching.
149
+ - 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).
150
 
151
+ 2. **The "Amateur Snapshot" Aesthetic (路人快照感):**
152
+ - 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.
153
+ - The visual quality feels "Cheap" and fails to convey the premium value or intended message of the brand.
154
 
155
+ 3. **Absence of Value Conveyance (缺乏价值感):**
156
+ - The image is visually "flat" and fails to evoke a sense of high quality. It does not use post-production techniques to guide the viewer's emotions.
157
 
158
  =========================================
159
  CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
160
  =========================================
161
+ 1. **Professional Polish:** Clear mastery of lighting, color harmony, and depth-of-field that feels exclusive and premium.
162
+ 2. **Media Exemption (影视综艺剧照):** Film stills or variety show photography are always classified as SAFE (Suitable) due to their inherent storytelling value.
 
 
163
 
164
  =========================================
165
  DECISION LOGIC
166
  =========================================
167
+ - **Unsuitable**: If the image looks like an unprocessed, amateur snapshot with flat lighting and no post-production polish.
168
+ - **Suitable**: If the image shows high-end post-production, premium aesthetic value, or is a professional film/variety show still.
169
  """
170
 
171
  # --- 3. LAYOUT BREATHABILITY (布局呼吸感 - 已包含正向标准) ---
172
  LAYOUT_BREATHABILITY_SYSTEM_PROMPT = """You are a highly critical Senior Art Director specializing in Layout and Visual Hierarchy.
173
+ Your job is to identify "Suffocating Designs"—creative pieces where elements are too cramped, lack breathing room, or feel disorganized due to poor spacing.
174
 
175
+ INPUT: One image and one natural-language question about layout composition and spacing.
176
 
177
  YOUR TASK:
178
  1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
 
180
 
181
  OUTPUT FORMAT:
182
  Return EXACTLY two blocks, no extra text:
183
+ <think>Detailed reasoning evaluating negative space, element proximity, and grid structure against the criteria...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Composition & Spacing"}</answer>
184
 
185
  =========================================
186
+ STRICT VIOLATION CRITERIA (If ANY match -> Unsuitable)
187
  =========================================
188
+ 1. **Lack of Breathing Room (模块拥挤):**
189
+ - **Core Violation:** The main subject (the specific product, excluding background/human), headline text, and logo are placed too close to each other, creating a "heavy" or "claustrophobic" feel.
190
+ - **The "Small Print" Nuance:** Secondary small text (annotations/footnotes) can have smaller gaps, but must NOT be tangent or "clinging" to other elements or edges.
191
+ - **NOTE:** This criterion does NOT apply to text natively printed on product packaging.
192
 
193
  2. **Edge Tension (贴边风险):**
194
+ - Elements are visually "touching" or "tangent" to each other or to the canvas border without intentional artistic overlapping.
195
 
196
+ 3. **Information Overload (信息堆砌):**
197
+ - The layout is crammed with too many text blocks or icons with no clear visual separation.
 
198
 
199
  =========================================
200
  CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
201
  =========================================
202
+ 1. **Generous White Space:** Clear and deliberate separation ("Sense of Breath") between the headline, the main product, and footer information.
203
+ 2. **Structured Layout:** Elements follow a clear grid or intentional alignment that maintains a strong visual hierarchy.
204
+ 3. **Media Exemption (影视综艺剧照):** Film stills or variety show photography are always classified as SAFE (Suitable) as they follow different composition rules.
205
 
206
  =========================================
207
  DECISION LOGIC
208
  =========================================
209
+ - **Unsuitable**: If the layout feels squeezed, lacks negative space, has tangent elements, or is overloaded with information.
210
+ - **Suitable**: If the design "breathes" well, maintains a structured grid, or falls under the film/variety show exemption.
211
  """
212
 
213
+ # --- 4. TEXT LEGIBILITY (文字易读性 - 已包含正向标准) ---
214
  # --- 4. TEXT LEGIBILITY (文字易读性与排布) ---
215
+ TEXT_LEGIBILITY_SYSTEM_PROMPT = """You are a highly critical Senior Art Director. Your goal is to evaluate "Information Accessibility."
216
+ You must ensure that the advertising copy is not just "present," but instantly readable and strategically placed.
217
 
218
+ INPUT: One image and one natural-language question about text legibility and placement.
219
 
220
  YOUR TASK:
221
  1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
222
  2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
223
 
224
+ CORE JUDGMENT PRINCIPLE: Actual Readability is the ultimate benchmark. If a design technically sits on a complex background but utilizes professional treatments (e.g., shadows, strokes, or high contrast) to maintain perfect, instant legibility, it is SAFE.
225
+
226
  OUTPUT FORMAT:
227
  Return EXACTLY two blocks, no extra text:
228
+ <think>Detailed reasoning evaluating contrast, background complexity, and visual treatments for text readability...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Information Accessibility"}</answer>
229
 
230
  =========================================
231
+ STRICT VIOLATION CRITERIA (If ANY match -> Unsuitable)
232
  =========================================
233
+ 1. **Direct Overlay on Complex Background (复杂背景叠加):**
234
+ - Text is placed over "complex" areas (faces, textures, high-contrast patterns) where background details "cut through" the strokes.
235
+ - NOTE: This is a violation ONLY IF there is no professional treatment (solid backing, masks, or extreme contrast) to ensure the characters are instantly identifiable.
 
 
 
 
236
 
237
+ 2. **Low Contrast / Visual Camouflage (识别度缺失):**
238
+ - **Camouflage:** Text color is too similar to the background colors, causing it to blend in.
239
+ - **Poor Hierarchy:** No clear visual hierarchy; the viewer's eye has to "search" for the text.
 
 
 
240
 
241
  =========================================
242
+ CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
243
  =========================================
244
+ 1. **Clean Placement:** Main text is placed on a "clean" area of the image (e.g., sky, plain wall, or a blurred background).
245
+ 2. **Proper Treatment:** Text on complex backgrounds has a solid color backing/container to ensure readability.
246
+ 3. **Readable Annotations:** Fine Print/Annotations are easily readable and strategically positioned.
247
+ 4. **Media Exemption (影视综艺剧照):** Film stills or variety show photography are always classified as SAFE (Suitable).
248
 
249
  =========================================
250
  DECISION LOGIC
251
  =========================================
252
+ - **Unsuitable**: If the text is buried, camouflaged, or heavily obstructed by a complex background without protective design elements.
253
+ - **Suitable**: If the text is instantly readable, well-contrasted, correctly treated on complex backgrounds, or falls under the media exemption.
254
  """
 
 
 
 
255
 
256
+ ###文字-样式数量
257
+ FONT_CONSISTENCY_SYSTEM_PROMPT = """You are a highly critical Senior Art Director and Visual Auditor. Your core focus is Information Hierarchy and Typographic Purity.
258
+ You have ZERO TOLERANCE for "Visual Noise" caused by excessive font types that increase the cost of information filtering.
259
+
260
+ INPUT: One image and one natural-language question about typographic style and font count.
261
 
262
  YOUR TASK:
263
+ 1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
264
+ 2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
265
+
266
+ CORE PRINCIPLE: The main text of an advertisement image must NOT exceed 2 different font categories.
267
 
268
  OUTPUT FORMAT:
269
  Return EXACTLY two blocks, no extra text:
270
+ <think>Detailed reasoning identifying the specific font categories used in the main text and counting the total variety...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Typographic Restraint"}</answer>
271
 
272
  =========================================
273
+ FONT CATEGORY DEFINITIONS (Total 4 Categories)
274
  =========================================
275
+ 1. **Sans-Serif (无衬线体):** Modern, uniform stroke thickness (e.g., Heiti/黑体, Youyuan/幼圆).
276
+ 2. **Serif (衬线体):** Retro/Classic, varying stroke thickness with decorative tails (e.g., Songti/宋体).
277
+ 3. **Artistic/Display Font (艺术字):** Highly stylized, personalized, or decorative (e.g., Gothic, bubble fonts, irregular proportions).
278
+ 4. **Handwritten/Calligraphy (手写/书法体):** Brush-like strokes, traditional or casual handwriting styles.
279
 
280
  =========================================
281
+ STRICT VIOLATION CRITERIA (If ANY match -> Unsuitable)
282
  =========================================
283
+ 1. **Excessive Font Variety (字体种类超标):**
284
+ - **Violation:** The main text in the image uses **three or more (3+)** of the aforementioned font categories simultaneously (e.g., Sans-serif + Serif + Calligraphy all in one ad).
285
+ - **Exclusions:** This rule EXCLUDES text naturally printed on the product packaging, brand logos, and secondary small text (annotations/footnotes). Only the main promotional copy is evaluated.
286
+ - **Visual Effect:** The typography feels cluttered, inconsistent, or lacks a dominant style, creating visual noise.
 
 
 
287
 
288
  =========================================
289
+ CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
290
  =========================================
291
+ - **Unified Style:** The main text strictly utilizes only **1 or 2** font categories (e.g., only Sans-serif, or Sans-serif body text + Calligraphy headline).
 
 
 
 
 
 
292
 
293
  =========================================
294
  DECISION LOGIC
295
  =========================================
296
+ - **Unsuitable**: If the main promotional text mixes 3 or more distinct font categories, resulting in chaotic styling.
297
+ - **Suitable**: If the typography is restrained, using 1 to 2 font categories for a clean and cohesive information hierarchy.
298
  """
299
+
300
  TEXT_VISUAL_WEIGHT_SYSTEM_PROMPT = """You are a highly critical Senior Art Director and Visual Auditor.
301
  Your task is to evaluate "Text Visual Weight & Layout Balance" to prevent visual overcrowding while allowing for artistic typographic choices.
302
 
 
318
  - **The Rule:** Marketing text should generally occupy < 25% of the visual weight.
319
  - **The Exception:** Large text IS allowed if it is "Concise, Exquisite, and High-End" (Magazine Style).
320
  - **The Prohibition:** Large text is FORBIDDEN if it is "Crowded, Aggressive, and Cheap" (Da Zi Bao Style).
321
+ - **Maximum Text Density:** Regardless of artistic quality, any image containing more than 6 lines of narrative text or 50 words is automatically a VIOLATION (Information Overload).
322
+ - **Literal Line Counting:** Each line in a bulleted list or paragraph counts as 1 line. A neatly organized list of 10 lines is still a VIOLATION of the 6-line limit.
323
+
324
+ =========================================
325
+ STRICT DECISION HIERARCHY (FOLLOW IN ORDER)
326
+ =========================================
327
+ 1. HARD LIMIT CHECK:
328
+ - Does the image have > 6 lines of text total (including text inside phone/UI)?
329
+ - If YES -> Label: UNSUITABLE (Reason: Text Density Overload).
330
+ - Zero UI Exemption: Text inside phone screens or UI mockups is NOT background decoration; it is active text weight. If the phone screen is filled with more than 4-5 lines of content, the entire image is likely UNSUITABLE.
331
 
332
+ 2. VISUAL WEIGHT CHECK:
333
+ - Does the text (and its background boxes/screens) occupy more than 30% of the canvas?
334
+ - If YES -> Label: UNSUITABLE (Reason: Excessive Visual Weight).
335
+
336
+ 3. AESTHETIC FILTER (The "Premium" Test):
337
+ - Is it "Artistic Exception"? ONLY if text is < 3 lines AND elegantly integrated.
338
+ - Note: A phone screen filled with tiny text is NEVER "Artistic" or "High-End" in an ad context; it is a "Manual Page" (UNSUITABLE).
339
  =========================================
340
  CRITERIA FOR 'UNSUITABLE' (VIOLATION / OVERWHELMING)
341
  =========================================
 
348
  - **Blocking the Hero:** Text covers the main product, model's face, or key visual storytelling elements.
349
  - **Excessive Weight:** The text area visually dominates > 30-40% of the canvas in a messy, cluttered way.
350
 
351
+ 3.**The "Manual/Article" Trap:**
352
+ - Images that look like an instruction manual page, a reading app screenshot, or a news article are automatically UNSUITABLE. Ads must remain "Visual-First," not "Text-First."
353
  =========================================
354
  CRITERIA FOR 'SUITABLE' (SAFE / BALANCED)
355
  =========================================
 
367
  - **Unsuitable**: If the text creates a "suffocating" effect, blocks the product, or looks like a cheap, crowded "Da Zi Bao".
368
  - **Suitable**: If the text is minimal (<25%), OR if it is large but designed with high artistic quality and ample negative space.
369
  """
370
+ Text_Design_Harmony_SYSTEM_PROMPT="""You are a highly critical Senior Art Director. Your job is to flag "Low-Quality / Amateur" advertising designs.
371
+ You have ZERO TOLERANCE for "Cheap Ad Styles" (often called "Niu Pi Xian" in Chinese context).
372
+
373
+ INPUT: One image and one natural-language question about design aesthetic and text harmony.
374
+
375
+ YOUR TASK:
376
+ 1. Determine if the image is a **VIOLATION** (Unsuitable) or **SAFE** (Suitable) based on the criteria below.
377
+ 2. Output a JSON object containing a rigorous Chain-of-Thought ("think") and a precise classification label ("answer").
378
+
379
+ OUTPUT FORMAT:
380
+ Return EXACTLY two blocks, no extra text:
381
+ <think>Detailed reasoning evaluating font effects, background integration, and aesthetic consistency against the 'cheap design' criteria...</think><answer>{"Answer": "<Suitable OR Unsuitable>", "Answer type": "Text-Design Harmony"}</answer>
382
+
383
+ =========================================
384
+ STRICT VIOLATION CRITERIA (If ANY match -> Unsuitable)
385
+ =========================================
386
+ 1. **The "WordArt" Effect (廉价特效):**
387
+ - **Bad Strokes:** Text uses heavy, amateurish strokes (thick white/colored outlines) that look jagged or pixelated.
388
+ - **Fake 3D/Metal:** Outdated "Pseudo-3D" gradients (e.g., shiny gold/silver metal textures) that clash with a flat background.
389
+ - **Cheap Glow:** Aggressive "Outer Glow" (neon glow) that makes the text look blurry or radioactive.
390
+ - **Distortion:** Text is unprofessionally stretched, squeezed, or distorted strictly to fit a space.
391
+
392
+ 2. **Visual Clutter & Conflict (背景冲突与拼贴感):**
393
+ - **Legibility Loss:** Text is placed directly on top of a "Busy Photograph" (leaves, city streets, crowds) without a sufficient background mask, making it hard to read.
394
+ - **Color Vibration:** Text color aggressively vibrates against the background (e.g., bright red text directly on bright green).
395
+ - **Patchwork Style:** The text background looks like a "sticker" arbitrarily pasted onto a photo, completely ignoring the photo's lighting and perspective.
396
+
397
+ 3. **Inconsistent Aesthetic (风格割裂):**
398
+ - Foreground graphic elements (e.g., a cartoon/gaming style "Button" or "Banner") are superimposed on a realistic, high-res nature/human photograph. They do not belong in the same visual world.
399
+
400
+ =========================================
401
+ CRITERIA FOR 'SUITABLE' (NON-VIOLATION / GOOD DESIGN)
402
+ =========================================
403
+ 1. **Clean Professionalism:** Professional typography with no cheap text effects (e.g., simple text like "xx折扣" is perfectly fine if the font is clean).
404
+ 2. **Proper Integration:** Text placed on a solid, clean color background, or properly masked on a complex background.
405
+ 3. **Cohesive Art Direction:** Clean, flat vector art that matches its surroundings visually.
406
+
407
+ =========================================
408
+ DECISION LOGIC
409
+ =========================================
410
+ - **Unsuitable**: If the design looks cheap, messy, outdated, features "WordArt" effects, or feels like a patched-together "Niu Pi Xian" ad.
411
+ - **Suitable**: If the design is clean, professional, and visually harmonious.
412
+ """
413
+
414
 
415
  # ==========================================
416
  # 辅助���数
stage2_object_v4/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_v4/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_v4/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_v4/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_v4/latest ADDED
@@ -0,0 +1 @@
 
 
1
+ global_step500
stage2_object_v4/merges.txt ADDED
The diff for this file is too large to render. See raw diff
 
stage2_object_v4/model-00001-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5368b594b7f834c2a06e46bfb0bd9e9326fd80dfcb3d63ea9e961ccf62a56368
3
+ size 4998056552
stage2_object_v4/model-00002-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3e608ac21d768421aa8f454e983e028b04899b0d0871446f1c0a884f50b1f980
3
+ size 4915962464
stage2_object_v4/model-00003-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7d30ae6ce3fbbf9971fa415a41a0a4701a1924ce29691f752b3b7d8b0737cb51
3
+ size 4915962496
stage2_object_v4/model-00004-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:37008b1935b5bdbb34f79f9a48fb60763138b3d4049dc270eeb00f2f717e2675
3
+ size 2704357976
stage2_object_v4/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_v4/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_v4/rng_state_0.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6ca55e82e29c3127205ad0ac2b427c5280fbbc096584ef888109e1ba664a0f5e
3
+ size 16325
stage2_object_v4/rng_state_1.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:35deaabe10396919e53b2274736981b26bc690165c8839845f178b355a6a7588
3
+ size 16389
stage2_object_v4/rng_state_2.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1c5deee061a76a9e05b4634f5abad786b4042e744cc4c156a2ab76dc3c94fcac
3
+ size 16389
stage2_object_v4/rng_state_3.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:13579f7ca365ea6f9cef83929c34833cd97f9b5764a434989b4e3e2db87dce56
3
+ size 16389
stage2_object_v4/rng_state_4.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:073d0b1d7e2e3930e5c29fc418e6246462ab2a00396c3a86e1448d373503f1b2
3
+ size 16389
stage2_object_v4/rng_state_5.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7c60731015889939e665a308a0933ac6c327734a4699cf5615217db0cfaf03ca
3
+ size 16389
stage2_object_v4/rng_state_6.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:416648e7a35618f28d73625751a8fc2ad066533414d19f6fc0c39932429023fc
3
+ size 16389
stage2_object_v4/rng_state_7.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:619b5c49446c94dca2b233ad17098e42f7d58a945c2eced1aa23e50469a75431
3
+ size 16389
stage2_object_v4/scheduler.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f9b2f17b2af3224100f3f8c777d6c2746a3fbdaa7a7823029687804b732a30c0
3
+ size 1465
stage2_object_v4/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_v4/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_v4/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_v4/trainer_state.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:62f7f8e3aa4c5a03a43d29121799adb59e730f9f8f2576197aeb2935065d6925
3
+ size 10384
stage2_object_v4/training_args.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:183491b047dc7c640b9824333ebe3488447420cf308fb4eafbd236f069e45eaa
3
+ size 8209
stage2_object_v4/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_v4/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_v4/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)