import os
import sys
import re
import gradio as gr
import torch
import numpy as np
import torch.nn.functional as F
# Matplotlib configuration for headless server rendering
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import matplotlib.animation as animation
from transformers import AutoModelForCausalLM, AutoTokenizer
from safetensors.torch import load_file
from huggingface_hub import snapshot_download
from threading import Thread
from transformers import TextIteratorStreamer
# ==========================================
# 1. Dynamic Environment Setup
# ==========================================
MODEL_ID = "Wojtekb30/Qwen2.5-1.5B-Instruct-RVQ-Human-Motion-CoT-PoC"
print("Fetching RVQ configuration and weights from repository...")
snapshot_download(
repo_id=MODEL_ID,
allow_patterns=["rvq_model/*"],
local_dir="./"
)
# Append to path so we can import the downloaded RVQ architecture
sys.path.append("./rvq_model")
from rvq_model import MotionRVQ_VAE
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# ==========================================
# 2. Load Models and Weights
# ==========================================
print("Loading LLM and Tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
llm_model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
device_map="auto"
)
llm_model.eval()
print("Loading RVQ decoder...")
rvq_model = MotionRVQ_VAE().to(device)
state_dict = load_file("./rvq_model/motion_rvq_weights.safetensors", device=str(device))
rvq_model.load_state_dict(state_dict)
rvq_model.eval()
mean = np.load('./rvq_model/Mean.npy')
std = np.load('./rvq_model/Std.npy')
# ==========================================
# 3. 3D Joint Processing Function
# ==========================================
def get_3d_joints(data_263):
frames = data_263.shape[0]
joints = np.zeros((frames, 22, 3))
for i in range(frames):
root_y = data_263[i, 3]
# Root joint position
joints[i, 0] = [0, root_y, 0]
# Convert local joint coordinates into world-space positions
local_positions = data_263[i, 4:67].reshape(21, 3)
joints[i, 1:] = local_positions + [0, root_y, 0]
return joints
# ==========================================
# 4. Rendering Pipeline
# ==========================================
def render_animation(all_tokens, output_path="output_animation.mp4"):
num_frames = len(all_tokens) // 4
token_matrix = np.zeros((4, num_frames), dtype=np.int64)
for i in range(num_frames):
for lvl in range(4):
token_idx = i * 4 + lvl
if token_idx < len(all_tokens):
parsed_lvl, val = all_tokens[token_idx]
token_matrix[lvl, i] = int(val)
token_tensor = torch.tensor(token_matrix, device=device).unsqueeze(0)
# Decode tokens into 3D motion
with torch.no_grad():
z_q = 0
for lvl in range(4):
indices = token_tensor[:, lvl, :]
quantizer = rvq_model.rvq.quantizers[lvl]
level_z_q = F.embedding(indices, quantizer.embedding)
level_z_q = level_z_q.permute(0, 2, 1)
z_q = z_q + level_z_q
reconstructed_motion = rvq_model.decoder(z_q)
recon_data = reconstructed_motion.squeeze(0).permute(1, 0).cpu().numpy()
recon_data = (recon_data * std) + mean
T_frames = recon_data.shape[0]
joints_recon = get_3d_joints(recon_data)
kinematic_tree = [
[0, 1, 4, 7, 10],
[0, 2, 5, 8, 11],
[0, 3, 6, 9, 12, 15],
[9, 13, 16, 18, 20],
[9, 14, 17, 19, 21]
]
fig = plt.figure(figsize=(6, 6))
ax = fig.add_subplot(111, projection='3d')
def update(frame):
ax.clear()
ax.set_title(f"Generated Motion\nFrame: {frame}/{T_frames}")
ax.set_xlim(-1, 1)
ax.set_ylim(-1, 1)
ax.set_zlim(0, 2)
ax.view_init(elev=10., azim=-90)
ax.axis('off')
for chain in kinematic_tree:
ax.plot(
joints_recon[frame, chain, 0],
joints_recon[frame, chain, 2],
joints_recon[frame, chain, 1],
linewidth=3,
marker='o',
markersize=4,
color='red'
)
ani = animation.FuncAnimation(fig, update, frames=T_frames, interval=50, repeat=False)
# CRITICAL: Save as an mp4 instead of calling plt.show()
ani.save(output_path, writer='ffmpeg', fps=20)
plt.close(fig) # Prevent memory leaks on the server
# ==========================================
# 5. Core Generation Logic (STREAMING)
# ==========================================
def generate_motion(user_prompt):
chat = [
{"role": "system", "content": "You are an embodied AI. You reason about your physical state and output precise motor actions inside tags."},
{"role": "user", "content": user_prompt}
]
input_text = tokenizer.apply_chat_template(chat, tokenize=False, add_generation_prompt=True)
inputs = tokenizer([input_text], return_tensors="pt").to(device)
# 1. Set up the Streamer
# skip_special_tokens=False ensures we still get our and tags!
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=False)
# 2. Run generation in a background thread so it doesn't block the UI
generation_kwargs = dict(
**inputs,
streamer=streamer,
max_new_tokens=512,
temperature=0.5,
do_sample=True
)
thread = Thread(target=llm_model.generate, kwargs=generation_kwargs)
thread.start()
# 3. Stream the text to the UI in real-time
generated_text = ""
for new_token in streamer:
generated_text += new_token
# Clean up the output for the display
display_text = generated_text.replace("<|im_end|>", "").strip()
# Yield the current text, and 'None' for the video (since it's not ready yet)
yield display_text, None
# 4. Once streaming is finished, parse the tokens and render the video
final_response = generated_text.replace("<|im_end|>", "").strip()
move_blocks = re.findall(r'(.*?)', final_response, re.DOTALL)
all_tokens = []
if move_blocks:
for block in move_blocks:
tokens = re.findall(r'', block)
all_tokens.extend(tokens)
video_path = None
if all_tokens:
video_path = "output_animation.mp4"
render_animation(all_tokens, video_path)
else:
final_response += "\n\nâ ď¸ [System Warning: The model did not output a valid sequence of tokens for rendering.]"
# 5. Final UI update: Provide the complete text AND the rendered video
yield final_response, video_path
# ==========================================
# 6. Gradio Interface
# ==========================================
with gr.Blocks(title="Qwen Human Motion CoT", theme=gr.themes.Soft()) as demo:
gr.Markdown("# đśââď¸ Qwen2.5-1.5B Human Motion Generator")
gr.Markdown("Type an instruction below. The model will generate a Chain-of-Thought planning its physical state, followed by discrete RVQ tokens decoded into 3D motion.\n\nThis space is a proof of concept based on: https://huggingface.co/Wojtekb30/Qwen2.5-1.5B-Instruct-RVQ-Human-Motion-CoT-PoC")
with gr.Row():
with gr.Column():
prompt_input = gr.Textbox(
label="Action Instruction",
lines=3,
placeholder="Walk forward."
)
generate_btn = gr.Button("Generate Animation", variant="primary")
with gr.Column():
video_output = gr.Video(label="Rendered 3D Skeleton")
with gr.Row():
text_output = gr.Textbox(label="Reasoning Chain & Generated Tokens", lines=12)
generate_btn.click(
fn=generate_motion,
inputs=prompt_input,
outputs=[text_output, video_output]
)
if __name__ == "__main__":
demo.launch()