This commit is contained in:
Jasonzzt
2026-03-12 05:33:13 +00:00
parent b3d8f97ae9
commit 1d2e76264a
2 changed files with 19 additions and 4 deletions
+14 -1
View File
@@ -43,7 +43,8 @@ def load_model_once():
checkpoint_path = "/PATH/TO/OmniLottie"
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
device = torch.device("cuda:0" if torch.cuda.is_available() else "xpu:0" if torch.xpu.is_available() else "cpu")
print(f"Loading model from {checkpoint_path}...")
model = LottieDecoder(pix_len=4560, text_len=1500)
@@ -599,6 +600,8 @@ def process_text_to_lottie(text_prompt, max_tokens, use_sampling, temperature, t
del inputs
if torch.cuda.is_available():
torch.cuda.empty_cache()
if torch.xpu.is_available():
torch.xpu.empty_cache()
lottie_json = tokens_to_lottie_json(generated_ids)
@@ -615,6 +618,8 @@ def process_text_to_lottie(text_prompt, max_tokens, use_sampling, temperature, t
except Exception as e:
if torch.cuda.is_available():
torch.cuda.empty_cache()
if torch.xpu.is_available():
torch.xpu.empty_cache()
return None, f"❌ Error: {str(e)}", None
def load_image_from_file(file_path):
@@ -663,6 +668,8 @@ def process_image_to_lottie(image_file, text_description, max_tokens, use_sampli
del inputs
if torch.cuda.is_available():
torch.cuda.empty_cache()
if torch.xpu.is_available():
torch.xpu.empty_cache()
lottie_json = tokens_to_lottie_json(generated_ids)
@@ -679,6 +686,8 @@ def process_image_to_lottie(image_file, text_description, max_tokens, use_sampli
except Exception as e:
if torch.cuda.is_available():
torch.cuda.empty_cache()
if torch.xpu.is_available():
torch.xpu.empty_cache()
return None, f"❌ Error: {str(e)}", None
def process_video_to_lottie(video, max_tokens, use_sampling, temperature, top_p, top_k):
@@ -709,6 +718,8 @@ def process_video_to_lottie(video, max_tokens, use_sampling, temperature, top_p,
del inputs
if torch.cuda.is_available():
torch.cuda.empty_cache()
if torch.xpu.is_available():
torch.xpu.empty_cache()
lottie_json = tokens_to_lottie_json(generated_ids)
@@ -725,6 +736,8 @@ def process_video_to_lottie(video, max_tokens, use_sampling, temperature, top_p,
except Exception as e:
if torch.cuda.is_available():
torch.cuda.empty_cache()
if torch.xpu.is_available():
torch.xpu.empty_cache()
return None, f"❌ Error: {str(e)}", None
def create_gradio_interface():
+5 -3
View File
@@ -811,7 +811,7 @@ def run_batch_text_file_inference(args, cfg):
"""
Generate Lottie from text file.
"""
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
device = torch.device("cuda:0" if torch.cuda.is_available() else "xpu:0" if torch.xpu.is_available() else "cpu")
print(f"Using device: {device}")
if not os.path.exists(args.batch_text_file):
@@ -918,7 +918,7 @@ def run_mmlottie_bench_inference(args, cfg):
- Task types: Text-to-Lottie, Text-Image-to-Lottie, Video-to-Lottie
- Fields: id, text, image, video, task_type, subset, etc.
"""
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
device = torch.device("cuda:0" if torch.cuda.is_available() else "xpu:0" if torch.xpu.is_available() else "cpu")
print(f"Using device: {device}")
# 1. Load dataset
@@ -1165,6 +1165,8 @@ def run_mmlottie_bench_inference(args, cfg):
# Clear cache
if torch.cuda.is_available():
torch.cuda.empty_cache()
if torch.xpu.is_available():
torch.xpu.empty_cache()
except Exception as e:
print(f" ❌ Error: {e}")
@@ -1186,7 +1188,7 @@ def run_mmlottie_bench_inference(args, cfg):
def run_single_inference(args, cfg):
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
device = torch.device("cuda:0" if torch.cuda.is_available() else "xpu:0" if torch.xpu.is_available() else "cpu")
print("Loading model...")
processor = AutoProcessor.from_pretrained(cfg['tokenizer_name'], padding_side="left")