优惠卷计算逻辑
This commit is contained in:
282
src/test/java/com/cst/video/TestVideo1.java
Normal file
282
src/test/java/com/cst/video/TestVideo1.java
Normal file
@@ -0,0 +1,282 @@
|
||||
package com.cst.video;
|
||||
|
||||
import com.alibaba.fastjson.JSON;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.http.*;
|
||||
import org.springframework.web.client.ResourceAccessException;
|
||||
import org.springframework.web.client.RestTemplate;
|
||||
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.nio.file.Paths;
|
||||
import java.time.LocalDateTime;
|
||||
import java.time.format.DateTimeFormatter;
|
||||
import java.util.Base64;
|
||||
import java.util.HashMap;
|
||||
import java.util.Map;
|
||||
|
||||
public class TestVideo1 {
|
||||
|
||||
private static final Logger log = LoggerFactory.getLogger(TestVideo1.class);
|
||||
private final RestTemplate restTemplate = new RestTemplate();
|
||||
|
||||
public static void main(String[] args) {
|
||||
System.out.println("\n" + "=".repeat(60));
|
||||
System.out.println("LTX-Video 视频生成客户端");
|
||||
System.out.println("=".repeat(60));
|
||||
/**
|
||||
* width height 是否有效
|
||||
* 704 480 ✓
|
||||
* 640 480 ✓
|
||||
* 1024 576 ✓
|
||||
*/
|
||||
String host = "192.168.1.35";
|
||||
int port = 8004;
|
||||
String prompt = "A beautiful young woman with feminine char🐻2pm wearing a bikini, dancing energetically on a sunny beach with golden sand and blue ocean waves, cinematic lighting, high quality, realistic";
|
||||
String negativePrompt = "ugly, deformed, blurry, low quality, pixelated, cartoon, anime, watermark, text, logo, multiple people, man, old woman";
|
||||
int numFrames = 240;
|
||||
int height = 480;
|
||||
int width = 704;
|
||||
int numInferenceSteps = 50;
|
||||
double guidanceScale = 7.5;
|
||||
int fps = 24;
|
||||
Integer seed = null;
|
||||
String outputDir = "/home/lizh/python_env/gen_video";
|
||||
|
||||
for (int i = 0; i < args.length; i++) {
|
||||
switch (args[i]) {
|
||||
case "--host":
|
||||
host = args[++i];
|
||||
break;
|
||||
case "--port":
|
||||
port = Integer.parseInt(args[++i]);
|
||||
break;
|
||||
case "--prompt":
|
||||
prompt = args[++i];
|
||||
break;
|
||||
case "--negative-prompt":
|
||||
// 负提示词参数,暂不使用
|
||||
++i;
|
||||
break;
|
||||
case "--num-frames":
|
||||
numFrames = Integer.parseInt(args[++i]);
|
||||
break;
|
||||
case "--height":
|
||||
height = Integer.parseInt(args[++i]);
|
||||
break;
|
||||
case "--width":
|
||||
width = Integer.parseInt(args[++i]);
|
||||
break;
|
||||
case "--num-inference-steps":
|
||||
numInferenceSteps = Integer.parseInt(args[++i]);
|
||||
break;
|
||||
case "--guidance-scale":
|
||||
guidanceScale = Double.parseDouble(args[++i]);
|
||||
break;
|
||||
case "--fps":
|
||||
fps = Integer.parseInt(args[++i]);
|
||||
break;
|
||||
case "--seed":
|
||||
seed = Integer.parseInt(args[++i]);
|
||||
break;
|
||||
case "--output-dir":
|
||||
outputDir = args[++i];
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
String baseUrl = "http://" + host + ":" + port;
|
||||
System.out.println("服务地址: " + baseUrl);
|
||||
System.out.println("=".repeat(60));
|
||||
|
||||
TestVideo1 client = new TestVideo1();
|
||||
|
||||
if (!client.healthCheck(baseUrl)) {
|
||||
System.err.println("\n✗ 服务不可用,请检查服务状态");
|
||||
System.exit(1);
|
||||
}
|
||||
|
||||
if (prompt == null || prompt.isEmpty()) {
|
||||
System.err.println("\n✗ 请提供 --prompt 参数");
|
||||
System.exit(1);
|
||||
}
|
||||
|
||||
if (height % 32 != 0) {
|
||||
System.err.println("\n✗ height 必须是 32 的倍数,当前值: " + height);
|
||||
System.exit(1);
|
||||
}
|
||||
if (width % 32 != 0) {
|
||||
System.err.println("\n✗ width 必须是 32 的倍数,当前值: " + width);
|
||||
System.exit(1);
|
||||
}
|
||||
|
||||
String result = client.generateVideo(
|
||||
baseUrl,
|
||||
prompt,
|
||||
negativePrompt,
|
||||
numFrames,
|
||||
height,
|
||||
width,
|
||||
numInferenceSteps,
|
||||
guidanceScale,
|
||||
fps,
|
||||
seed,
|
||||
outputDir
|
||||
);
|
||||
|
||||
if (result != null) {
|
||||
System.exit(0);
|
||||
} else {
|
||||
System.exit(1);
|
||||
}
|
||||
}
|
||||
|
||||
public boolean healthCheck(String baseUrl) {
|
||||
try {
|
||||
String url = baseUrl + "/health";
|
||||
ResponseEntity<String> resp = restTemplate.getForEntity(url, String.class);
|
||||
|
||||
if (resp.getStatusCode() == HttpStatus.OK) {
|
||||
JSONObject data = JSON.parseObject(resp.getBody());
|
||||
System.out.println("✓ 服务健康检查通过");
|
||||
System.out.println(" 状态: " + data.getString("status"));
|
||||
System.out.println(" 模型: " + data.getString("model"));
|
||||
System.out.println(" GPU: " + data.getString("gpu"));
|
||||
return true;
|
||||
} else {
|
||||
System.err.println("✗ 服务响应异常: HTTP " + resp.getStatusCode());
|
||||
return false;
|
||||
}
|
||||
} catch (ResourceAccessException e) {
|
||||
System.err.println("✗ 服务连接失败: " + e.getMessage());
|
||||
return false;
|
||||
} catch (Exception e) {
|
||||
System.err.println("✗ 服务健康检查失败: " + e.getMessage());
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
public void listModels(String baseUrl) {
|
||||
try {
|
||||
String url = baseUrl + "/v1/models";
|
||||
ResponseEntity<String> resp = restTemplate.getForEntity(url, String.class);
|
||||
|
||||
if (resp.getStatusCode() == HttpStatus.OK) {
|
||||
JSONObject data = JSON.parseObject(resp.getBody());
|
||||
System.out.println("可用模型列表:");
|
||||
for (Object modelObj : data.getJSONArray("data")) {
|
||||
JSONObject model = (JSONObject) modelObj;
|
||||
System.out.println(" - " + model.getString("id"));
|
||||
}
|
||||
} else {
|
||||
System.err.println("✗ 获取模型列表失败: HTTP " + resp.getStatusCode());
|
||||
}
|
||||
} catch (Exception e) {
|
||||
System.err.println("✗ 获取模型列表失败: " + e.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
public String generateVideo(
|
||||
String baseUrl,
|
||||
String prompt,
|
||||
String negativePrompt,
|
||||
int numFrames,
|
||||
int height,
|
||||
int width,
|
||||
int numInferenceSteps,
|
||||
double guidanceScale,
|
||||
int fps,
|
||||
Integer seed,
|
||||
String outputDir
|
||||
) {
|
||||
Map<String, Object> payload = new HashMap<>();
|
||||
payload.put("prompt", prompt);
|
||||
payload.put("negative_prompt", negativePrompt);
|
||||
payload.put("num_frames", numFrames);
|
||||
payload.put("height", height);
|
||||
payload.put("width", width);
|
||||
payload.put("num_inference_steps", numInferenceSteps);
|
||||
payload.put("guidance_scale", guidanceScale);
|
||||
payload.put("fps", fps);
|
||||
if (seed != null) {
|
||||
payload.put("seed", seed);
|
||||
}
|
||||
|
||||
System.out.println("\n" + "=".repeat(60));
|
||||
System.out.println("开始生成视频...");
|
||||
System.out.println("=".repeat(60));
|
||||
System.out.println("提示词: " + (prompt.length() > 60 ? prompt.substring(0, 60) + "..." : prompt));
|
||||
System.out.println("负提示词: " + (negativePrompt.isEmpty() ? "无" : (negativePrompt.length() > 60 ? negativePrompt.substring(0, 60) + "..." : negativePrompt)));
|
||||
System.out.println("帧数: " + numFrames + ", 分辨率: " + width + "x" + height);
|
||||
System.out.println("推理步数: " + numInferenceSteps + ", 引导系数: " + guidanceScale);
|
||||
System.out.println("FPS: " + fps + ", 种子: " + (seed != null ? seed : "随机"));
|
||||
System.out.println("=".repeat(60));
|
||||
|
||||
long startTime = System.currentTimeMillis();
|
||||
try {
|
||||
String url = baseUrl + "/v1/video/generate";
|
||||
|
||||
HttpHeaders headers = new HttpHeaders();
|
||||
headers.setContentType(MediaType.APPLICATION_JSON);
|
||||
|
||||
HttpEntity<Map<String, Object>> entity = new HttpEntity<>(payload, headers);
|
||||
|
||||
ResponseEntity<String> resp = restTemplate.exchange(
|
||||
url,
|
||||
HttpMethod.POST,
|
||||
entity,
|
||||
String.class
|
||||
);
|
||||
|
||||
long elapsed = System.currentTimeMillis() - startTime;
|
||||
|
||||
if (resp.getStatusCode() == HttpStatus.OK) {
|
||||
JSONObject data = JSON.parseObject(resp.getBody());
|
||||
String videoBase64 = data.getString("video_base64");
|
||||
Integer seedUsed = data.getInteger("seed");
|
||||
Double duration = data.getDouble("duration_seconds");
|
||||
|
||||
Path outputPath = Paths.get(outputDir);
|
||||
if (!Files.exists(outputPath)) {
|
||||
Files.createDirectories(outputPath);
|
||||
}
|
||||
|
||||
String timestamp = LocalDateTime.now().format(DateTimeFormatter.ofPattern("yyyyMMdd_HHmmss"));
|
||||
String filename = "video_" + timestamp + "_seed" + seedUsed + ".mp4";
|
||||
Path filepath = outputPath.resolve(filename);
|
||||
|
||||
byte[] videoBytes = Base64.getDecoder().decode(videoBase64);
|
||||
Files.write(filepath, videoBytes);
|
||||
|
||||
System.out.println("\n✓ 视频生成成功!");
|
||||
System.out.println(" 耗时: " + (elapsed / 1000.0) + " 秒");
|
||||
System.out.println(" 使用种子: " + seedUsed);
|
||||
System.out.println(" 视频时长: " + String.format("%.2f", duration) + " 秒");
|
||||
System.out.println(" 保存路径: " + filepath.toAbsolutePath());
|
||||
return filepath.toString();
|
||||
} else {
|
||||
long elapsedTime = System.currentTimeMillis() - startTime;
|
||||
System.err.println("\n✗ 视频生成失败: HTTP " + resp.getStatusCode());
|
||||
try {
|
||||
JSONObject errorData = JSON.parseObject(resp.getBody());
|
||||
System.err.println(" 错误信息: " + errorData.getString("detail"));
|
||||
} catch (Exception e) {
|
||||
System.err.println(" 响应内容: " + resp.getBody());
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
} catch (ResourceAccessException e) {
|
||||
long elapsedTime = System.currentTimeMillis() - startTime;
|
||||
System.err.println("\n✗ 请求超时 (" + (elapsedTime / 1000.0) + " 秒)");
|
||||
return null;
|
||||
} catch (Exception e) {
|
||||
long elapsedTime = System.currentTimeMillis() - startTime;
|
||||
System.err.println("\n✗ 请求失败 (" + (elapsedTime / 1000.0) + " 秒): " + e.getMessage());
|
||||
e.printStackTrace();
|
||||
return null;
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user