#!/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()