lightx2v's picture
Upload folder using huggingface_hub
a8a33c5 verified
Raw
History Blame Contribute Delete
5.48 kB
#!/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()