tomkay commited on
Commit
057b9e5
·
verified ·
1 Parent(s): 26612fd

Upload generate.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. generate.py +172 -0
generate.py ADDED
@@ -0,0 +1,172 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Generate video with a RAM-quantized LTX-2.3 model via ltx-2-mlx.
2
+
3
+ The script loads the mixed-precision model directory produced by
4
+ reformat_ltx_for_pipeline.py into dgrauet/ltx-2-mlx's DistilledPipeline,
5
+ replacing the default apply_quantization with a per-layer mixed-precision
6
+ version that handles our variable-bits allocation.
7
+
8
+ Usage:
9
+ python experiments/flux_phase1/generate_ltx.py \\
10
+ --model-dir results/ltx-2.3/model_dir_12gb/ \\
11
+ --prompt "A cat walks through a field of flowers" \\
12
+ --output output.mp4
13
+
14
+ # Specify resolution and frame count
15
+ python experiments/flux_phase1/generate_ltx.py \\
16
+ --model-dir results/ltx-2.3/model_dir_12gb/ \\
17
+ --prompt "Ocean waves at sunset" \\
18
+ --height 480 --width 704 --num-frames 97 \\
19
+ --output ocean.mp4
20
+
21
+ # Quiet (no progress output)
22
+ python experiments/flux_phase1/generate_ltx.py \\
23
+ --model-dir results/ltx-2.3/model_dir_12gb/ \\
24
+ --prompt "..." --output out.mp4 --quiet
25
+ """
26
+ import argparse
27
+ import sys
28
+ from collections import defaultdict
29
+ from pathlib import Path
30
+
31
+ import mlx.core as mx
32
+ import mlx.nn as nn
33
+
34
+
35
+ def apply_mixed_precision_quantization(
36
+ model: nn.Module,
37
+ weights: dict,
38
+ group_size: int = 64,
39
+ ) -> None:
40
+ """Per-layer mixed-precision quantization from a weight dict.
41
+
42
+ Unlike ltx_core_mlx's apply_quantization (which uses a single detected
43
+ bit width for all layers), this version detects each layer's bits from
44
+ its packed weight shape and applies nn.quantize once per unique bit width.
45
+
46
+ Layers that have .scales but whose bits can't be determined are skipped
47
+ (kept as nn.Linear — they will fail at load_weights if shapes mismatch,
48
+ which surfaces any genuine key errors).
49
+ """
50
+ layer_bits: dict[str, int] = {}
51
+ for key in weights:
52
+ if not key.endswith(".scales"):
53
+ continue
54
+ layer = key[: -len(".scales")]
55
+ w_key = layer + ".weight"
56
+ if w_key not in weights:
57
+ continue
58
+ w_cols = weights[w_key].shape[-1]
59
+ s_cols = weights[key].shape[-1]
60
+ bits = round(w_cols * 32 / (s_cols * group_size))
61
+ if bits in (2, 3, 4, 5, 6, 8):
62
+ layer_bits[layer] = bits
63
+
64
+ if not layer_bits:
65
+ return
66
+
67
+ bits_to_layers: dict[int, set] = defaultdict(set)
68
+ for layer, b in layer_bits.items():
69
+ bits_to_layers[b].add(layer)
70
+
71
+ for bits, layers in sorted(bits_to_layers.items()):
72
+ def _predicate(path: str, module: nn.Module, _layers=layers) -> bool:
73
+ return path in _layers and isinstance(module, nn.Linear)
74
+ nn.quantize(model, group_size=group_size, bits=bits, class_predicate=_predicate)
75
+
76
+ total = sum(len(v) for v in bits_to_layers.values())
77
+ dist = {b: len(v) for b, v in sorted(bits_to_layers.items())}
78
+ print(f" Mixed-precision quantization: {total} layers — {dist}", flush=True)
79
+
80
+
81
+ def _patch_pipeline_quantization(group_size: int = 64):
82
+ """Monkeypatch ltx_core_mlx to use our mixed-precision quantizer."""
83
+ import ltx_core_mlx.utils.weights as wm
84
+ import ltx_pipelines_mlx.utils._orchestration as orch
85
+
86
+ def _patched_apply(model, weights, group_size=group_size, bits=None):
87
+ apply_mixed_precision_quantization(model, weights, group_size)
88
+
89
+ wm.apply_quantization = _patched_apply
90
+ # Also patch the reference in _orchestration (it imports apply_quantization
91
+ # at module level in some builds)
92
+ if hasattr(orch, "apply_quantization"):
93
+ orch.apply_quantization = _patched_apply
94
+
95
+
96
+ def main():
97
+ p = argparse.ArgumentParser(
98
+ description="Generate video with RAM-quantized LTX-2.3."
99
+ )
100
+ p.add_argument("--model-dir", required=True,
101
+ help="Model directory from reformat_ltx_for_pipeline.py")
102
+ p.add_argument("--prompt", required=True, help="Text prompt for video generation")
103
+ p.add_argument("--output", default="output.mp4", help="Output video path")
104
+ p.add_argument("--height", type=int, default=480)
105
+ p.add_argument("--width", type=int, default=704)
106
+ p.add_argument("--num-frames", type=int, default=97)
107
+ p.add_argument("--frame-rate", type=float, default=24.0)
108
+ p.add_argument("--seed", type=int, default=42)
109
+ p.add_argument("--stage1-steps", type=int, default=None)
110
+ p.add_argument("--stage2-steps", type=int, default=None)
111
+ p.add_argument("--gemma-model", default="mlx-community/gemma-3-12b-it-4bit",
112
+ help="Gemma model ID for text encoding")
113
+ p.add_argument("--low-memory", action="store_true", default=True,
114
+ help="Aggressive memory management (default: on)")
115
+ p.add_argument("--no-low-memory", dest="low_memory", action="store_false")
116
+ p.add_argument("--quiet", action="store_true", help="Suppress pipeline progress output")
117
+ args = p.parse_args()
118
+
119
+ model_dir = Path(args.model_dir)
120
+ if not model_dir.exists():
121
+ print(f"Error: model dir {model_dir} not found", file=sys.stderr)
122
+ sys.exit(1)
123
+
124
+ required = [
125
+ "transformer-distilled.safetensors",
126
+ "connector.safetensors",
127
+ "vae_decoder.safetensors",
128
+ "audio_vae.safetensors",
129
+ "vocoder.safetensors",
130
+ ]
131
+ missing = [f for f in required if not (model_dir / f).exists()]
132
+ if missing:
133
+ print(f"Error: missing files in {model_dir}: {missing}", file=sys.stderr)
134
+ print("Run reformat_ltx_for_pipeline.py first.", file=sys.stderr)
135
+ sys.exit(1)
136
+
137
+ # Patch before importing the pipeline
138
+ print("Patching quantization to use per-layer mixed-precision…")
139
+ _patch_pipeline_quantization()
140
+
141
+ from ltx_pipelines_mlx.distilled import DistilledPipeline
142
+
143
+ print(f"Loading pipeline from {model_dir}…")
144
+ pipeline = DistilledPipeline(
145
+ model_dir=str(model_dir),
146
+ gemma_model_id=args.gemma_model,
147
+ low_memory=args.low_memory,
148
+ )
149
+
150
+ print(f"\nGenerating: '{args.prompt}'")
151
+ print(f" Resolution: {args.height}×{args.width}, {args.num_frames} frames @ {args.frame_rate} fps")
152
+
153
+ video_latent, audio_latent = pipeline.generate_two_stage(
154
+ prompt=args.prompt,
155
+ height=args.height,
156
+ width=args.width,
157
+ num_frames=args.num_frames,
158
+ frame_rate=args.frame_rate,
159
+ seed=args.seed,
160
+ stage1_steps=args.stage1_steps,
161
+ stage2_steps=args.stage2_steps,
162
+ )
163
+
164
+ print(f"\nDecoding and saving → {args.output}")
165
+ out = pipeline._decode_and_save_video(
166
+ video_latent, audio_latent, args.output, frame_rate=args.frame_rate
167
+ )
168
+ print(f"Done: {out}")
169
+
170
+
171
+ if __name__ == "__main__":
172
+ main()