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()