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