muragekibicho commited on
Commit
2071f4a
·
verified ·
1 Parent(s): 300a3ee

Upload 3 files

Browse files
converted_safetensors/model_epoch_25.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8452569c153034ebfe2288dbbed9dddf3a0ef0a46def244d5cd46b499f17113b
3
+ size 46964372
model_epoch_25.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8873b6e5de2a6982011b5600bd121b563f2ce54deb9fbda7e637bf1d402b7c91
3
+ size 46983138
safetensors_converter.py ADDED
@@ -0,0 +1,864 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Convert PyTorch model files (.pt/.pth) to .safetensors."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import importlib.metadata
7
+ import json
8
+ import logging
9
+ import sys
10
+ import time
11
+ import warnings
12
+ from collections import Counter
13
+ from dataclasses import dataclass
14
+ from datetime import datetime
15
+ from pathlib import Path
16
+ from typing import Any, Mapping
17
+
18
+ import torch
19
+ from colorama import Fore, Style, init
20
+ from packaging import version
21
+ from safetensors.torch import load_file as load_safetensors_file
22
+ from safetensors.torch import save_file
23
+
24
+
25
+ warnings.filterwarnings("ignore")
26
+ logging.getLogger("torch").setLevel(logging.ERROR)
27
+
28
+ MIN_SAFETENSORS_VERSION = "0.4.1"
29
+ SUPPORTED_EXTENSIONS = (".pt", ".pth")
30
+
31
+
32
+ @dataclass(frozen=True)
33
+ class ConversionResult:
34
+ status: str
35
+ reason_code: str
36
+ message: str
37
+ input_file: str
38
+ output_file: str
39
+ validation_warnings: list[str]
40
+
41
+
42
+ @dataclass
43
+ class RuntimeOptions:
44
+ allow_unsafe_load: bool
45
+ ask_unsafe_once: bool
46
+ cast_float32: bool
47
+ verbose: bool
48
+ validate: bool
49
+ strict_validate: bool
50
+ json_report: bool
51
+ dry_run: bool
52
+
53
+
54
+ @dataclass(frozen=True)
55
+ class ValidationOutcome:
56
+ success: bool
57
+ errors: list[str]
58
+ warnings: list[str]
59
+
60
+
61
+ def parse_args(argv: list[str]) -> argparse.Namespace:
62
+ parser = argparse.ArgumentParser(
63
+ description=(
64
+ "Converts PyTorch model files (.pt/.pth) to .safetensors. "
65
+ "Original files are never modified."
66
+ )
67
+ )
68
+ parser.add_argument(
69
+ "input_path",
70
+ help="Single model file, or folder containing .pt/.pth files",
71
+ )
72
+ parser.add_argument(
73
+ "output_dir",
74
+ nargs="?",
75
+ default=None,
76
+ help=(
77
+ "Output folder for converted files. "
78
+ "Default: converted_safetensors inside the input folder"
79
+ ),
80
+ )
81
+ parser.add_argument(
82
+ "--verbose",
83
+ action="store_true",
84
+ help="Print extra details while processing",
85
+ )
86
+ parser.add_argument(
87
+ "--allow-unsafe-load",
88
+ action="store_true",
89
+ help=(
90
+ "If weights_only=True loading fails, retry with weights_only=False "
91
+ "without asking each time"
92
+ ),
93
+ )
94
+ parser.add_argument(
95
+ "--cast-float32",
96
+ action="store_true",
97
+ help=(
98
+ "Cast floating tensors to float32 before saving (can improve compatibility with picky loaders/tools, but may increase file size and reduce precision)"
99
+ ),
100
+ )
101
+ parser.add_argument(
102
+ "--skip-validate",
103
+ action="store_true",
104
+ help="Skip post-conversion validation checks",
105
+ )
106
+ parser.add_argument(
107
+ "--strict-validate",
108
+ action="store_true",
109
+ help="Fail validation on any key/shape/dtype mismatch",
110
+ )
111
+ parser.add_argument(
112
+ "--json-report",
113
+ action="store_true",
114
+ help="Write a detailed JSON report to the output folder",
115
+ )
116
+ parser.add_argument(
117
+ "--dry-run",
118
+ action="store_true",
119
+ help=(
120
+ "Show what would be converted without loading model files or writing outputs"
121
+ ),
122
+ )
123
+
124
+ args = parser.parse_args(argv[1:])
125
+ if args.strict_validate and args.skip_validate:
126
+ parser.error("--strict-validate cannot be combined with --skip-validate")
127
+ return args
128
+
129
+
130
+ def resolve_paths(input_path_raw: str, output_dir_raw: str | None) -> tuple[Path, Path]:
131
+ input_path = Path(input_path_raw).expanduser().resolve()
132
+ if not input_path.exists():
133
+ raise FileNotFoundError(f"Input path does not exist: {input_path}")
134
+
135
+ if output_dir_raw:
136
+ output_dir = Path(output_dir_raw).expanduser().resolve()
137
+ elif input_path.is_file():
138
+ output_dir = input_path.parent / "converted_safetensors"
139
+ else:
140
+ output_dir = input_path / "converted_safetensors"
141
+
142
+ return input_path, output_dir
143
+
144
+
145
+ def check_safetensors_support() -> bool | None:
146
+ installed = importlib.metadata.version("safetensors")
147
+ supports_large_files = version.parse(installed) >= version.parse(
148
+ MIN_SAFETENSORS_VERSION
149
+ )
150
+
151
+ if supports_large_files:
152
+ print(
153
+ Fore.GREEN + f"* safetensors {installed} supports files larger than 4 GB."
154
+ )
155
+ return True
156
+
157
+ print(
158
+ Fore.RED
159
+ + Style.BRIGHT
160
+ + "\n* Warning *\n"
161
+ + Style.NORMAL
162
+ + (
163
+ f"Installed safetensors version ({installed}) can only handle models under 4 GB.\n"
164
+ "Larger models will be skipped."
165
+ )
166
+ )
167
+ user_input = (
168
+ input(
169
+ Fore.YELLOW
170
+ + f"Continue with safetensors {installed} and skip >4 GB models? y/[n] :: "
171
+ )
172
+ .strip()
173
+ .lower()
174
+ )
175
+ if user_input != "y":
176
+ print(
177
+ Fore.YELLOW
178
+ + (
179
+ f"\nUpdate safetensors to {MIN_SAFETENSORS_VERSION} or newer and run again.\n"
180
+ "Exiting ...\n"
181
+ )
182
+ )
183
+ return None
184
+
185
+ print(Fore.YELLOW + "\n* Continuing with legacy size limitation active.\n")
186
+ return False
187
+
188
+
189
+ def collect_input_files(input_path: Path) -> list[Path]:
190
+ if input_path.is_file():
191
+ return [input_path]
192
+
193
+ return [
194
+ p
195
+ for p in sorted(input_path.iterdir())
196
+ if p.is_file() and p.suffix.lower() in SUPPORTED_EXTENSIONS
197
+ ]
198
+
199
+
200
+ def get_state_dict(checkpoint: Any) -> Any:
201
+ if isinstance(checkpoint, torch.nn.Module):
202
+ return checkpoint.state_dict()
203
+ if isinstance(checkpoint, Mapping):
204
+ return checkpoint.get("state_dict", checkpoint)
205
+ return checkpoint
206
+
207
+
208
+ def extract_tensors(obj: Any, prefix: str = "") -> dict[str, torch.Tensor]:
209
+ tensors: dict[str, torch.Tensor] = {}
210
+
211
+ if isinstance(obj, torch.Tensor):
212
+ key = prefix or "tensor"
213
+ tensors[key] = obj
214
+ return tensors
215
+
216
+ if isinstance(obj, Mapping):
217
+ for key, value in obj.items():
218
+ key_str = str(key)
219
+ next_prefix = f"{prefix}.{key_str}" if prefix else key_str
220
+ tensors.update(extract_tensors(value, next_prefix))
221
+ return tensors
222
+
223
+ if isinstance(obj, (list, tuple)):
224
+ for idx, value in enumerate(obj):
225
+ next_prefix = f"{prefix}.{idx}" if prefix else str(idx)
226
+ tensors.update(extract_tensors(value, next_prefix))
227
+ return tensors
228
+
229
+ return tensors
230
+
231
+
232
+ def prepare_tensors(
233
+ tensors: dict[str, torch.Tensor], cast_float32: bool
234
+ ) -> dict[str, torch.Tensor]:
235
+ prepared: dict[str, torch.Tensor] = {}
236
+ for name, tensor in tensors.items():
237
+ t = tensor.detach().cpu().contiguous()
238
+ if cast_float32 and t.is_floating_point():
239
+ t = t.float()
240
+ prepared[name] = t
241
+ return prepared
242
+
243
+
244
+ def tensor_bytes(tensors: Mapping[str, torch.Tensor]) -> int:
245
+ return sum(t.numel() * t.element_size() for t in tensors.values())
246
+
247
+
248
+ def bytes_to_gb(size_bytes: int) -> float:
249
+ return size_bytes / (1024**3)
250
+
251
+
252
+ def load_checkpoint(
253
+ input_file: Path,
254
+ runtime: RuntimeOptions,
255
+ unsafe_retry_enabled: bool,
256
+ ) -> tuple[Any | None, bool, str | None]:
257
+ try:
258
+ return (
259
+ torch.load(str(input_file), map_location="cpu", weights_only=True),
260
+ unsafe_retry_enabled,
261
+ None,
262
+ )
263
+ except TypeError:
264
+ try:
265
+ return (
266
+ torch.load(str(input_file), map_location="cpu"),
267
+ unsafe_retry_enabled,
268
+ None,
269
+ )
270
+ except Exception as exc:
271
+ return None, unsafe_retry_enabled, str(exc)
272
+ except Exception as exc:
273
+ initial_error = str(exc)
274
+
275
+ should_try_unsafe = runtime.allow_unsafe_load
276
+ if (
277
+ not should_try_unsafe
278
+ and runtime.ask_unsafe_once
279
+ and not unsafe_retry_enabled
280
+ and not runtime.dry_run
281
+ ):
282
+ print(
283
+ Fore.YELLOW
284
+ + Style.BRIGHT
285
+ + "\n* Info: safe load failed for this file.\n"
286
+ + Style.NORMAL
287
+ + (
288
+ "You can retry with weights_only=False, which can execute code inside the model file.\n"
289
+ "Use this only for trusted model sources."
290
+ )
291
+ )
292
+ user_input = (
293
+ input(
294
+ Fore.YELLOW
295
+ + "Retry with weights_only=False for this run? y/[n] :: "
296
+ )
297
+ .strip()
298
+ .lower()
299
+ )
300
+ should_try_unsafe = user_input == "y"
301
+ unsafe_retry_enabled = should_try_unsafe
302
+
303
+ if not should_try_unsafe:
304
+ return None, unsafe_retry_enabled, initial_error
305
+
306
+ try:
307
+ return (
308
+ torch.load(str(input_file), map_location="cpu", weights_only=False),
309
+ unsafe_retry_enabled,
310
+ None,
311
+ )
312
+ except TypeError:
313
+ try:
314
+ return (
315
+ torch.load(str(input_file), map_location="cpu"),
316
+ unsafe_retry_enabled,
317
+ None,
318
+ )
319
+ except Exception as unsafe_exc:
320
+ return (
321
+ None,
322
+ unsafe_retry_enabled,
323
+ "Safe load failed and unsafe retry also failed:\n"
324
+ + f"weights_only=True error:\n{initial_error}\n"
325
+ + f"weights_only=False error:\n{unsafe_exc}",
326
+ )
327
+ except Exception as unsafe_exc:
328
+ return (
329
+ None,
330
+ unsafe_retry_enabled,
331
+ "Safe load failed and unsafe retry also failed:\n"
332
+ + f"weights_only=True error:\n{initial_error}\n"
333
+ + f"weights_only=False error:\n{unsafe_exc}",
334
+ )
335
+
336
+
337
+ def validate_saved_output(
338
+ prepared_tensors: dict[str, torch.Tensor],
339
+ output_file: Path,
340
+ strict: bool,
341
+ ) -> ValidationOutcome:
342
+ try:
343
+ saved_tensors = load_safetensors_file(str(output_file), device="cpu")
344
+ except Exception as exc:
345
+ return ValidationOutcome(
346
+ success=False,
347
+ errors=[f"Could not load produced safetensors file: {exc}"],
348
+ warnings=[],
349
+ )
350
+
351
+ src_keys = set(prepared_tensors.keys())
352
+ dst_keys = set(saved_tensors.keys())
353
+
354
+ missing_keys = sorted(src_keys - dst_keys)
355
+ extra_keys = sorted(dst_keys - src_keys)
356
+
357
+ errors: list[str] = []
358
+ warnings: list[str] = []
359
+
360
+ if missing_keys:
361
+ errors.append(f"Missing {len(missing_keys)} tensor key(s) in output")
362
+ if extra_keys:
363
+ errors.append(f"Found {len(extra_keys)} extra tensor key(s) in output")
364
+
365
+ common = sorted(src_keys & dst_keys)
366
+ shape_mismatches = 0
367
+ dtype_mismatches = 0
368
+
369
+ for key in common:
370
+ src = prepared_tensors[key]
371
+ dst = saved_tensors[key]
372
+ if tuple(src.shape) != tuple(dst.shape):
373
+ shape_mismatches += 1
374
+ if src.dtype != dst.dtype:
375
+ dtype_mismatches += 1
376
+
377
+ if shape_mismatches > 0:
378
+ errors.append(f"{shape_mismatches} tensor shape mismatch(es)")
379
+
380
+ if dtype_mismatches > 0:
381
+ if strict:
382
+ errors.append(f"{dtype_mismatches} tensor dtype mismatch(es)")
383
+ else:
384
+ warnings.append(f"{dtype_mismatches} tensor dtype mismatch(es)")
385
+
386
+ return ValidationOutcome(
387
+ success=len(errors) == 0,
388
+ errors=errors,
389
+ warnings=warnings,
390
+ )
391
+
392
+
393
+ def plan_dry_run(
394
+ input_file: Path,
395
+ output_dir: Path,
396
+ supports_large_files: bool,
397
+ ) -> ConversionResult:
398
+ output_file = output_dir / f"{input_file.stem}.safetensors"
399
+
400
+ notes: list[str] = []
401
+ reason_code = "DRYRUN_READY"
402
+
403
+ if output_file.exists():
404
+ reason_code = "DRYRUN_OVERWRITE"
405
+ notes.append("Output already exists and would be overwritten")
406
+
407
+ if not supports_large_files:
408
+ file_size = input_file.stat().st_size
409
+ if file_size >= 4 * (1024**3):
410
+ reason_code = "DRYRUN_SIZE_RISK"
411
+ notes.append(
412
+ "Input file itself is >= 4 GB, and current safetensors may fail depending on tensor payload size"
413
+ )
414
+
415
+ message = (
416
+ "Dry run: conversion not executed"
417
+ if not notes
418
+ else "Dry run: " + "; ".join(notes)
419
+ )
420
+
421
+ return ConversionResult(
422
+ status="DRYRUN",
423
+ reason_code=reason_code,
424
+ message=message,
425
+ input_file=str(input_file),
426
+ output_file=str(output_file),
427
+ validation_warnings=[],
428
+ )
429
+
430
+
431
+ def convert_file(
432
+ input_file: Path,
433
+ output_dir: Path,
434
+ supports_large_files: bool,
435
+ runtime: RuntimeOptions,
436
+ unsafe_retry_enabled: bool,
437
+ ) -> tuple[ConversionResult, bool]:
438
+ if runtime.dry_run:
439
+ return plan_dry_run(
440
+ input_file, output_dir, supports_large_files
441
+ ), unsafe_retry_enabled
442
+
443
+ checkpoint, unsafe_retry_enabled, load_error = load_checkpoint(
444
+ input_file, runtime, unsafe_retry_enabled
445
+ )
446
+ if load_error is not None:
447
+ return (
448
+ ConversionResult(
449
+ status="FAILED",
450
+ reason_code="FAIL_LOAD",
451
+ message=f"Could not load checkpoint:\n{load_error}",
452
+ input_file=str(input_file),
453
+ output_file="",
454
+ validation_warnings=[],
455
+ ),
456
+ unsafe_retry_enabled,
457
+ )
458
+
459
+ state_like = get_state_dict(checkpoint)
460
+ tensors = extract_tensors(state_like)
461
+ if not tensors:
462
+ return (
463
+ ConversionResult(
464
+ status="FAILED",
465
+ reason_code="FAIL_NO_TENSORS",
466
+ message="No tensors found in loaded object. Unsupported checkpoint structure.",
467
+ input_file=str(input_file),
468
+ output_file="",
469
+ validation_warnings=[],
470
+ ),
471
+ unsafe_retry_enabled,
472
+ )
473
+
474
+ prepared = prepare_tensors(tensors, cast_float32=runtime.cast_float32)
475
+
476
+ if not supports_large_files:
477
+ size = tensor_bytes(prepared)
478
+ if size >= 4 * (1024**3):
479
+ return (
480
+ ConversionResult(
481
+ status="SKIPPED",
482
+ reason_code="SKIP_SIZE_LIMIT",
483
+ message=(
484
+ "Tensor payload exceeds 4 GB "
485
+ f"({bytes_to_gb(size):.2f} GB). Current safetensors version cannot save it."
486
+ ),
487
+ input_file=str(input_file),
488
+ output_file="",
489
+ validation_warnings=[],
490
+ ),
491
+ unsafe_retry_enabled,
492
+ )
493
+
494
+ output_file = output_dir / f"{input_file.stem}.safetensors"
495
+
496
+ try:
497
+ save_file(prepared, str(output_file))
498
+ except Exception as exc:
499
+ err = str(exc)
500
+ if "invalid load key" in err.lower():
501
+ return (
502
+ ConversionResult(
503
+ status="FAILED",
504
+ reason_code="FAIL_INVALID_FORMAT",
505
+ message=f"Invalid/corrupted input or unsupported format:\n{err}",
506
+ input_file=str(input_file),
507
+ output_file="",
508
+ validation_warnings=[],
509
+ ),
510
+ unsafe_retry_enabled,
511
+ )
512
+ if "non contiguous" in err.lower():
513
+ return (
514
+ ConversionResult(
515
+ status="FAILED",
516
+ reason_code="FAIL_NONCONTIG",
517
+ message=f"Failed after preparing contiguous tensors:\n{err}",
518
+ input_file=str(input_file),
519
+ output_file="",
520
+ validation_warnings=[],
521
+ ),
522
+ unsafe_retry_enabled,
523
+ )
524
+ return (
525
+ ConversionResult(
526
+ status="FAILED",
527
+ reason_code="FAIL_SAVE",
528
+ message=f"Save failed for an unexpected reason:\n{err}",
529
+ input_file=str(input_file),
530
+ output_file="",
531
+ validation_warnings=[],
532
+ ),
533
+ unsafe_retry_enabled,
534
+ )
535
+
536
+ if runtime.validate:
537
+ validation = validate_saved_output(
538
+ prepared_tensors=prepared,
539
+ output_file=output_file,
540
+ strict=runtime.strict_validate,
541
+ )
542
+ if not validation.success:
543
+ return (
544
+ ConversionResult(
545
+ status="FAILED",
546
+ reason_code="FAIL_VALIDATE",
547
+ message="Validation failed: " + "; ".join(validation.errors),
548
+ input_file=str(input_file),
549
+ output_file=str(output_file),
550
+ validation_warnings=validation.warnings,
551
+ ),
552
+ unsafe_retry_enabled,
553
+ )
554
+
555
+ if validation.warnings:
556
+ return (
557
+ ConversionResult(
558
+ status="OK",
559
+ reason_code="OK_VALIDATED_WARN",
560
+ message="Converted and validated with warnings",
561
+ input_file=str(input_file),
562
+ output_file=str(output_file),
563
+ validation_warnings=validation.warnings,
564
+ ),
565
+ unsafe_retry_enabled,
566
+ )
567
+
568
+ return (
569
+ ConversionResult(
570
+ status="OK",
571
+ reason_code="OK_VALIDATED",
572
+ message="Converted and validated",
573
+ input_file=str(input_file),
574
+ output_file=str(output_file),
575
+ validation_warnings=[],
576
+ ),
577
+ unsafe_retry_enabled,
578
+ )
579
+
580
+ return (
581
+ ConversionResult(
582
+ status="OK",
583
+ reason_code="OK_NO_VALIDATE",
584
+ message="Converted (validation skipped)",
585
+ input_file=str(input_file),
586
+ output_file=str(output_file),
587
+ validation_warnings=[],
588
+ ),
589
+ unsafe_retry_enabled,
590
+ )
591
+
592
+
593
+ def print_file_status(
594
+ idx: int,
595
+ total: int,
596
+ model_file: Path,
597
+ result: ConversionResult,
598
+ verbose: bool,
599
+ ) -> None:
600
+ if result.status == "OK":
601
+ color = Fore.GREEN
602
+ elif result.status == "SKIPPED":
603
+ color = Fore.YELLOW
604
+ elif result.status == "DRYRUN":
605
+ color = Fore.CYAN
606
+ else:
607
+ color = Fore.RED
608
+
609
+ base = (
610
+ f"[{str(idx).zfill(3)}/{str(total).zfill(3)}] "
611
+ f"{model_file.name} -> {result.status} [{result.reason_code}]"
612
+ )
613
+ print(color + base)
614
+
615
+ if verbose:
616
+ print(Style.NORMAL + f" {result.message}")
617
+ for warning in result.validation_warnings:
618
+ print(Fore.YELLOW + f" validation warning: {warning}")
619
+
620
+
621
+ def print_final_report(
622
+ results: list[ConversionResult],
623
+ elapsed_seconds: float,
624
+ runtime: RuntimeOptions,
625
+ output_dir: Path,
626
+ json_path: Path | None,
627
+ ) -> None:
628
+ reason_legend = {
629
+ "OK_VALIDATED": "Converted and validation passed",
630
+ "OK_VALIDATED_WARN": "Converted and validated with warnings",
631
+ "OK_NO_VALIDATE": "Converted without validation",
632
+ "SKIP_SIZE_LIMIT": "Skipped due to legacy 4 GB limit",
633
+ "FAIL_LOAD": "Could not load checkpoint",
634
+ "FAIL_NO_TENSORS": "No tensors found in checkpoint",
635
+ "FAIL_INVALID_FORMAT": "Invalid/corrupted input format",
636
+ "FAIL_NONCONTIG": "Could not save after contiguous prep",
637
+ "FAIL_SAVE": "Save failed for another reason",
638
+ "FAIL_VALIDATE": "Post-save validation failed",
639
+ "DRYRUN_READY": "Dry run: ready to convert",
640
+ "DRYRUN_OVERWRITE": "Dry run: output exists and would be overwritten",
641
+ "DRYRUN_SIZE_RISK": "Dry run: potential size limitation risk",
642
+ }
643
+ failure_actions = {
644
+ "FAIL_LOAD": "Try --allow-unsafe-load only for trusted files; if still failing, verify file integrity/source.",
645
+ "FAIL_NO_TENSORS": "This file is likely not a plain tensor checkpoint; inspect its structure before converting.",
646
+ "FAIL_INVALID_FORMAT": "Check that the file is a valid .pt/.pth checkpoint and re-download if corruption is suspected.",
647
+ "FAIL_NONCONTIG": "Re-save the original checkpoint from PyTorch if possible, then retry conversion.",
648
+ "FAIL_SAVE": "Re-run with --verbose for detail and verify disk permissions/free space.",
649
+ "FAIL_VALIDATE": "Run again with --verbose and inspect key/shape/dtype mismatches before using the output.",
650
+ }
651
+
652
+ status_counts = Counter(r.status for r in results)
653
+ reason_counts = Counter(r.reason_code for r in results)
654
+
655
+ print(Fore.CYAN + Style.BRIGHT + "\n=| Run Report |=\n")
656
+
657
+ print(Fore.CYAN + Style.BRIGHT + "Summary")
658
+ print(Fore.CYAN + f"* Total files considered: {len(results)}")
659
+ print(Fore.GREEN + f"* OK: {status_counts.get('OK', 0)}")
660
+ print(Fore.YELLOW + f"* SKIPPED: {status_counts.get('SKIPPED', 0)}")
661
+ print(Fore.RED + f"* FAILED: {status_counts.get('FAILED', 0)}")
662
+ print(Fore.CYAN + f"* DRYRUN: {status_counts.get('DRYRUN', 0)}")
663
+ print(Fore.CYAN + f"* Validation: {'ON' if runtime.validate else 'OFF'}")
664
+ print(
665
+ Fore.CYAN
666
+ + f"* Validation strict mode: {'ON' if runtime.strict_validate else 'OFF'}"
667
+ )
668
+ print(Fore.CYAN + f"* Elapsed: {elapsed_seconds:.2f}s")
669
+
670
+ print(Fore.CYAN + Style.BRIGHT + "\nReason Code Breakdown")
671
+ for reason, count in sorted(reason_counts.items(), key=lambda item: item[0]):
672
+ print(Fore.CYAN + f"* {reason}: {count}")
673
+
674
+ print(Fore.CYAN + Style.BRIGHT + "\nReason Code Legend")
675
+ for reason in sorted(reason_counts.keys()):
676
+ description = reason_legend.get(reason, "No legend entry available")
677
+ print(Fore.CYAN + f"* {reason}: {description}")
678
+
679
+ failures = [r for r in results if r.status == "FAILED"]
680
+ skips = [r for r in results if r.status == "SKIPPED"]
681
+ warns = [r for r in results if r.validation_warnings]
682
+
683
+ if failures:
684
+ print(Fore.RED + Style.BRIGHT + "\nFailures")
685
+ for item in failures:
686
+ print(Fore.RED + f"* {Path(item.input_file).name}: {item.reason_code}")
687
+ action = failure_actions.get(
688
+ item.reason_code,
689
+ "Check verbose output and source checkpoint integrity, then retry.",
690
+ )
691
+ print(Fore.YELLOW + f" suggested action: {action}")
692
+ if runtime.verbose:
693
+ print(Fore.RED + f" {item.message}")
694
+
695
+ if skips:
696
+ print(Fore.YELLOW + Style.BRIGHT + "\nSkipped")
697
+ for item in skips:
698
+ print(Fore.YELLOW + f"* {Path(item.input_file).name}: {item.reason_code}")
699
+
700
+ if warns:
701
+ print(Fore.YELLOW + Style.BRIGHT + "\nValidation Warnings")
702
+ for item in warns:
703
+ print(Fore.YELLOW + f"* {Path(item.input_file).name}:")
704
+ for warning in item.validation_warnings:
705
+ print(Fore.YELLOW + f" - {warning}")
706
+
707
+ print(Fore.CYAN + Style.BRIGHT + "\nOutput")
708
+ print(Fore.CYAN + f"* Output folder: {output_dir}")
709
+ if json_path is not None:
710
+ print(Fore.CYAN + f"* JSON report: {json_path}")
711
+ else:
712
+ print(Fore.CYAN + "* JSON report: disabled")
713
+
714
+
715
+ def write_results_json(output_dir: Path, details: list[dict[str, Any]]) -> Path:
716
+ ts = datetime.now().strftime("%Y%m%d-%H%M%S")
717
+ path = output_dir / f"_results_{ts}.json"
718
+ with path.open("w", encoding="utf-8") as f:
719
+ json.dump(details, f, indent=2, ensure_ascii=False)
720
+ return path
721
+
722
+
723
+ def main(argv: list[str]) -> int:
724
+ init(autoreset=True)
725
+
726
+ print(
727
+ Fore.CYAN
728
+ + "\n "
729
+ + "-" * 39
730
+ + "\n--| "
731
+ + Style.BRIGHT
732
+ + "SafeTensors Converter Script"
733
+ + Style.NORMAL
734
+ + " |--\n"
735
+ + " "
736
+ + "-" * 39
737
+ + "\n"
738
+ )
739
+
740
+ args = parse_args(argv)
741
+
742
+ try:
743
+ input_path, output_dir = resolve_paths(args.input_path, args.output_dir)
744
+ except Exception as exc:
745
+ print(Fore.RED + Style.BRIGHT + "Error! " + Style.NORMAL + str(exc))
746
+ return 1
747
+
748
+ if not args.dry_run:
749
+ try:
750
+ output_dir.mkdir(parents=True, exist_ok=True)
751
+ except Exception as exc:
752
+ print(
753
+ Fore.RED
754
+ + Style.BRIGHT
755
+ + "Error! "
756
+ + Style.NORMAL
757
+ + f"Could not create output directory {output_dir}:\n{exc}"
758
+ )
759
+ return 1
760
+
761
+ if args.verbose:
762
+ print(Fore.CYAN + "* Checking installed safetensors version ...")
763
+
764
+ supports_large_files = check_safetensors_support()
765
+ if supports_large_files is None:
766
+ return 1
767
+
768
+ files = collect_input_files(input_path)
769
+ if not files:
770
+ print(Fore.YELLOW + "No .pt/.pth files found to process.")
771
+ return 0
772
+
773
+ runtime = RuntimeOptions(
774
+ allow_unsafe_load=args.allow_unsafe_load,
775
+ ask_unsafe_once=True,
776
+ cast_float32=args.cast_float32,
777
+ verbose=args.verbose,
778
+ validate=not args.skip_validate,
779
+ strict_validate=args.strict_validate,
780
+ json_report=args.json_report,
781
+ dry_run=args.dry_run,
782
+ )
783
+
784
+ if runtime.verbose:
785
+ mode = "file" if input_path.is_file() else "folder"
786
+ print(Fore.CYAN + "* Run configuration:")
787
+ print(Fore.CYAN + f" + mode: {mode}")
788
+ print(Fore.CYAN + f" + input: {input_path}")
789
+ print(Fore.CYAN + f" + output: {output_dir}")
790
+ print(Fore.CYAN + f" + files: {len(files)}")
791
+ print(Fore.CYAN + f" + cast float32: {runtime.cast_float32}")
792
+ print(Fore.CYAN + f" + validate: {runtime.validate}")
793
+ print(Fore.CYAN + f" + strict validate: {runtime.strict_validate}")
794
+ print(Fore.CYAN + f" + dry run: {runtime.dry_run}")
795
+ print(Fore.CYAN + f" + json report: {runtime.json_report}")
796
+
797
+ print(
798
+ Fore.CYAN
799
+ + (
800
+ f"\n* Planning {len(files)} model file(s) ..."
801
+ if runtime.dry_run
802
+ else f"\n* Processing {len(files)} model file(s) ..."
803
+ )
804
+ )
805
+
806
+ started = time.perf_counter()
807
+ unsafe_retry_enabled = runtime.allow_unsafe_load
808
+
809
+ all_results: list[ConversionResult] = []
810
+
811
+ for idx, model_file in enumerate(files, start=1):
812
+ result, unsafe_retry_enabled = convert_file(
813
+ model_file,
814
+ output_dir,
815
+ supports_large_files,
816
+ runtime,
817
+ unsafe_retry_enabled,
818
+ )
819
+ all_results.append(result)
820
+ print_file_status(idx, len(files), model_file, result, runtime.verbose)
821
+
822
+ elapsed = time.perf_counter() - started
823
+
824
+ json_path: Path | None = None
825
+ if runtime.json_report:
826
+ json_rows: list[dict[str, Any]] = []
827
+ for r in all_results:
828
+ json_rows.append({
829
+ "input_file": r.input_file,
830
+ "output_file": r.output_file,
831
+ "status": r.status,
832
+ "reason_code": r.reason_code,
833
+ "message": r.message,
834
+ "validation_warnings": r.validation_warnings,
835
+ })
836
+ if not runtime.dry_run:
837
+ json_path = write_results_json(output_dir, json_rows)
838
+ else:
839
+ # For dry runs, use a sibling report file without creating conversion output folders.
840
+ fallback_dir = output_dir if output_dir.exists() else input_path.parent
841
+ fallback_dir.mkdir(parents=True, exist_ok=True)
842
+ json_path = write_results_json(fallback_dir, json_rows)
843
+
844
+ print_final_report(
845
+ results=all_results,
846
+ elapsed_seconds=elapsed,
847
+ runtime=runtime,
848
+ output_dir=output_dir,
849
+ json_path=json_path,
850
+ )
851
+
852
+ if runtime.dry_run:
853
+ print(
854
+ Fore.CYAN + "\nDry run finished. No model files were loaded or converted.\n"
855
+ )
856
+ return 0
857
+
858
+ if any(r.status == "FAILED" for r in all_results):
859
+ return 2
860
+ return 0
861
+
862
+
863
+ if __name__ == "__main__":
864
+ sys.exit(main(sys.argv))