Skip to content

Navigation Menu

Sign in
Sign up

Add MiniMax M2.7 as alternative LLM provider for prompt refinement #10

New issue

Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.

By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.

Already on GitHub? Sign in to your account

Open
octo-patch wants to merge 1 commit into JavisVerse:main
base: main
Choose a base branch
Loading
from octo-patch:feature/add-minimax-provider
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 11 additions & 8 deletions gradio/app.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -175,10 +175,13 @@ def parse_args():
extract_json_from_prompts,
extract_prompts_loop,
get_random_prompt_by_openai,
get_random_prompt_by_llm,
has_openai_key,
has_minimax_key,
merge_prompt,
prepare_multi_resolution_info,
refine_prompts_by_openai,
refine_prompts_by_llm,
split_prompt,
)
from javisdit.utils.misc import to_torch_dtype
Expand Down Expand Up @@ -289,14 +292,14 @@ def run_inference(
batched_prompt_segment_list.append(prompt_segment_list)
batched_loop_idx_list.append(loop_idx_list)

# 1. refine prompt by openai
# 1. refine prompt by available LLM provider (MiniMax or OpenAI)
if refine_prompt:
# check if openai key is provided
if not has_openai_key():
gr.Warning("OpenAI API key is not provided, the prompt will not be enhanced.")
# check if any LLM key is provided
if not has_minimax_key() and not has_openai_key():
gr.Warning("No LLM API key provided (MINIMAX_API_KEY or OPENAI_API_KEY). The prompt will not be enhanced.")
else:
for idx, prompt_segment_list in enumerate(batched_prompt_segment_list):
batched_prompt_segment_list[idx] = refine_prompts_by_openai(prompt_segment_list)
batched_prompt_segment_list[idx] = refine_prompts_by_llm(prompt_segment_list)

# process scores
aesthetic_score = aesthetic_score if use_aesthetic_score else None
Expand Down Expand Up @@ -449,11 +452,11 @@ def run_video_inference(


def generate_random_prompt():
if "OPENAI_API_KEY" not in os.environ:
gr.Warning("Your prompt is empty and the OpenAI API key is not provided, please enter a valid prompt")
if not has_minimax_key() and "OPENAI_API_KEY" not in os.environ:
gr.Warning("Your prompt is empty and no LLM API key is provided (MINIMAX_API_KEY or OPENAI_API_KEY), please enter a valid prompt")
return None
else:
prompt_text = get_random_prompt_by_openai()
prompt_text = get_random_prompt_by_llm()
return prompt_text


Expand Down
82 changes: 80 additions & 2 deletions javisdit/utils/inference_utils.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -329,6 +329,7 @@ def dframe_to_frame(num):


OPENAI_CLIENT = None
MINIMAX_CLIENT = None
REFINE_PROMPTS = None
REFINE_PROMPTS_PATH = "assets/texts/t2v_pllava.txt"
REFINE_PROMPTS_TEMPLATE = """
Expand All @@ -345,6 +346,9 @@ def dframe_to_frame(num):
The prompt should pay attention to all objects in the video. The description should be useful for AI to re-generate the video. The description should be no more than six sentences. The prompt should be in English.
"""

MINIMAX_API_BASE = "https://api.minimax.io/v1"
MINIMAX_DEFAULT_MODEL = "MiniMax-M2.7"


def get_openai_response(sys_prompt, usr_prompt, model="gpt-4o"):
global OPENAI_CLIENT
Expand All @@ -370,6 +374,45 @@ def get_openai_response(sys_prompt, usr_prompt, model="gpt-4o"):
return completion.choices[0].message.content


def get_minimax_response(sys_prompt, usr_prompt, model=MINIMAX_DEFAULT_MODEL):
"""Call MiniMax LLM via OpenAI-compatible API for prompt refinement."""
global MINIMAX_CLIENT
if MINIMAX_CLIENT is None:
from openai import OpenAI

MINIMAX_CLIENT = OpenAI(
api_key=os.environ.get("MINIMAX_API_KEY"),
base_url=MINIMAX_API_BASE,
)

# MiniMax requires temperature in (0.0, 1.0]
completion = MINIMAX_CLIENT.chat.completions.create(
model=model,
temperature=0.7,
messages=[
{"role": "system", "content": sys_prompt},
{"role": "user", "content": usr_prompt},
],
)

return completion.choices[0].message.content


def has_openai_key():
return "OPENAI_API_KEY" in os.environ


def has_minimax_key():
return "MINIMAX_API_KEY" in os.environ


def get_llm_response(sys_prompt, usr_prompt):
"""Auto-detect available LLM provider: MiniMax takes priority over OpenAI."""
if has_minimax_key():
return get_minimax_response(sys_prompt, usr_prompt)
return get_openai_response(sys_prompt, usr_prompt)


def get_random_prompt_by_openai():
global RANDOM_PROMPTS
if RANDOM_PROMPTS is None:
Expand All @@ -390,8 +433,24 @@ def refine_prompt_by_openai(prompt):
return response


def has_openai_key():
return "OPENAI_API_KEY" in os.environ
def get_random_prompt_by_llm():
"""Generate a random video prompt using the available LLM provider."""
global RANDOM_PROMPTS
if RANDOM_PROMPTS is None:
examples = load_prompts(REFINE_PROMPTS_PATH)
RANDOM_PROMPTS = RANDOM_PROMPTS_TEMPLATE.format("\n".join(examples))

return get_llm_response(RANDOM_PROMPTS, "Generate one example.")


def refine_prompt_by_llm(prompt):
"""Refine a video generation prompt using the available LLM provider."""
global REFINE_PROMPTS
if REFINE_PROMPTS is None:
examples = load_prompts(REFINE_PROMPTS_PATH)
REFINE_PROMPTS = REFINE_PROMPTS_TEMPLATE.format("\n".join(examples))

return get_llm_response(REFINE_PROMPTS, prompt)


def refine_prompts_by_openai(prompts):
Expand All @@ -411,6 +470,25 @@ def refine_prompts_by_openai(prompts):
return new_prompts


def refine_prompts_by_llm(prompts):
"""Refine prompts using the available LLM provider (MiniMax or OpenAI)."""
provider = "MiniMax" if has_minimax_key() else "OpenAI"
new_prompts = []
for prompt in prompts:
try:
if prompt.strip() == "":
new_prompt = get_random_prompt_by_llm()
print(f"[Info] Empty prompt detected, generate random prompt via {provider}: {new_prompt}")
else:
new_prompt = refine_prompt_by_llm(prompt)
print(f"[Info] Refine prompt via {provider}: {prompt} -> {new_prompt}")
new_prompts.append(new_prompt)
except Exception as e:
print(f"[Warning] Failed to refine prompt via {provider}: {prompt} due to {e}")
new_prompts.append(prompt)
return new_prompts


def add_watermark(
input_video_path, watermark_image_path="./assets/images/watermark/watermark.png", output_video_path=None
):
Expand Down
Loading

AltStyle によって変換されたページ (->オリジナル) /