PEFT
Safetensors
English
lora
prompt-rewriting
minimax-h3
text-to-audio-video
audio-video-generation
Instructions to use lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3.6-27B") model = PeftModel.from_pretrained(base_model, "lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA") - Notebooks
- Google Colab
- Kaggle
| #!/usr/bin/env python3 | |
| """Rewrite one T2VA prompt with the MiniMax-H3 Prompt Rewriter LoRA.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import torch | |
| import transformers | |
| from transformers import AutoTokenizer, set_seed | |
| from prompt_template import build_messages | |
| DEFAULT_BASE_MODEL = "Qwen/Qwen3.6-27B" | |
| DEFAULT_ADAPTER = "lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA" | |
| RESOLUTIONS = ("21:9", "16:9", "4:3", "1:1", "3:4", "9:16") | |
| def get_model_class(): | |
| model_class = getattr(transformers, "AutoModelForImageTextToText", None) | |
| if model_class is None: | |
| model_class = getattr(transformers, "AutoModelForVision2Seq", None) | |
| if model_class is None: | |
| raise RuntimeError( | |
| "A recent Transformers version with AutoModelForImageTextToText " | |
| "support is required. Install the packages in requirements.txt." | |
| ) | |
| return model_class | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--prompt", required=True, help="Original text prompt to rewrite.") | |
| parser.add_argument("--duration", type=int, choices=range(4, 16), default=10, metavar="4..15") | |
| parser.add_argument("--resolution", choices=RESOLUTIONS, default="16:9", help="Target aspect ratio.") | |
| parser.add_argument("--base-model", default=DEFAULT_BASE_MODEL, help="HF model ID or local base-model path.") | |
| parser.add_argument("--adapter", default=DEFAULT_ADAPTER, help="HF model ID or local PEFT adapter path.") | |
| parser.add_argument( | |
| "--base-only", | |
| action="store_true", | |
| help="Skip the LoRA and run the Qwen base-model baseline.", | |
| ) | |
| parser.add_argument("--max-new-tokens", type=int, default=2048) | |
| parser.add_argument("--dtype", choices=("bfloat16", "float16", "float32"), default="bfloat16") | |
| parser.add_argument("--attn-implementation", choices=("sdpa", "flash_attention_2", "eager"), default="sdpa") | |
| parser.add_argument("--temperature", type=float, default=0.7) | |
| parser.add_argument("--top-p", type=float, default=0.8) | |
| parser.add_argument("--top-k", type=int, default=20) | |
| parser.add_argument("--repetition-penalty", type=float, default=1.05) | |
| parser.add_argument("--greedy", action="store_true", help="Use deterministic greedy decoding.") | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument( | |
| "--output", | |
| type=Path, | |
| help="Optional output path. A .json file stores conditions and output; other suffixes store plain text.", | |
| ) | |
| return parser.parse_args() | |
| def input_device(model: torch.nn.Module) -> torch.device: | |
| """Return the embedding device, including when Accelerate shards the model.""" | |
| embeddings = model.get_input_embeddings() | |
| if embeddings is not None and hasattr(embeddings, "weight"): | |
| return embeddings.weight.device | |
| return next(model.parameters()).device | |
| def main() -> None: | |
| args = parse_args() | |
| set_seed(args.seed) | |
| dtype = getattr(torch, args.dtype) | |
| tokenizer = AutoTokenizer.from_pretrained(args.base_model, trust_remote_code=True) | |
| if tokenizer.pad_token_id is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| model = get_model_class().from_pretrained( | |
| args.base_model, | |
| torch_dtype=dtype, | |
| device_map="auto", | |
| low_cpu_mem_usage=True, | |
| trust_remote_code=True, | |
| attn_implementation=args.attn_implementation, | |
| ) | |
| if not args.base_only: | |
| from peft import PeftModel | |
| model = PeftModel.from_pretrained(model, args.adapter) | |
| model.eval() | |
| rendered = tokenizer.apply_chat_template( | |
| build_messages(args.prompt, args.resolution, args.duration), | |
| tokenize=False, | |
| add_generation_prompt=True, | |
| enable_thinking=False, | |
| ) | |
| inputs = tokenizer(rendered, return_tensors="pt", add_special_tokens=False) | |
| device = input_device(model) | |
| inputs = {name: tensor.to(device) for name, tensor in inputs.items()} | |
| generation_kwargs = { | |
| "max_new_tokens": args.max_new_tokens, | |
| "do_sample": not args.greedy, | |
| "repetition_penalty": args.repetition_penalty, | |
| "pad_token_id": tokenizer.pad_token_id, | |
| "eos_token_id": tokenizer.eos_token_id, | |
| } | |
| if not args.greedy: | |
| generation_kwargs.update( | |
| temperature=args.temperature, | |
| top_p=args.top_p, | |
| top_k=args.top_k, | |
| ) | |
| with torch.inference_mode(): | |
| generated = model.generate(**inputs, **generation_kwargs) | |
| new_tokens = generated[:, inputs["input_ids"].shape[1] :] | |
| rewritten_prompt = tokenizer.batch_decode(new_tokens, skip_special_tokens=True)[0].strip() | |
| if args.output is not None: | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| if args.output.suffix.lower() == ".json": | |
| record = { | |
| "prompt": args.prompt.strip(), | |
| "resolution": args.resolution, | |
| "duration": args.duration, | |
| "rewritten_prompt": rewritten_prompt, | |
| "base_model": args.base_model, | |
| "adapter": None if args.base_only else args.adapter, | |
| } | |
| args.output.write_text(json.dumps(record, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") | |
| else: | |
| args.output.write_text(rewritten_prompt + "\n", encoding="utf-8") | |
| print(rewritten_prompt) | |
| if __name__ == "__main__": | |
| main() | |