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
File size: 5,479 Bytes
a8a33c5 | 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 | #!/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()
|