233 lines
8.5 KiB
Python
233 lines
8.5 KiB
Python
import gradio as gr
|
|
import time
|
|
from functools import partial
|
|
|
|
# 从各个模块导入所需的功能
|
|
import config
|
|
from analysis_functions import (
|
|
display_transfer_function,
|
|
time_domain_analysis,
|
|
frequency_domain_analysis,
|
|
root_locus_analysis
|
|
)
|
|
# ===== 新增:算例演示模块函数导入 =====
|
|
from case_demo_functions import run_case_demo
|
|
from chatbot import chat_with_ai
|
|
from user_stats import get_online_status_html, update_user_activity
|
|
from ui_components import (
|
|
create_header,
|
|
create_time_domain_tab,
|
|
create_frequency_domain_tab,
|
|
create_root_locus_tab,
|
|
create_case_demo_tab,
|
|
create_chatbot_tab
|
|
)
|
|
|
|
# 加载外部CSS文件
|
|
with open("assets/styles.css", "r", encoding="utf-8") as f:
|
|
custom_css = f.read()
|
|
|
|
# --- 主应用界面 ---
|
|
with gr.Blocks(title="自动控制理论学习网站 - AI+数智平台", css=custom_css) as demo:
|
|
# 1. 创建UI组件
|
|
# 用户会话ID(隐藏组件)
|
|
session_id = gr.State(value=lambda: str(time.time()) + "_" + str(hash(time.time())))
|
|
|
|
# 创建头部信息和在线计数器
|
|
online_counter = create_header()
|
|
|
|
# 创建共享的输入组件
|
|
with gr.Row():
|
|
with gr.Column(scale=1):
|
|
with gr.Group():
|
|
gr.HTML("<div class='card-title'>📊 通用系统参数</div>")
|
|
num_input = gr.Textbox(
|
|
label="传递函数分子系数 (Numerator)",
|
|
value="1",
|
|
placeholder="例如: 1 或 1,2,3",
|
|
info="💡 用逗号分隔多个系数,从最高次项到常数项"
|
|
)
|
|
den_input = gr.Textbox(
|
|
label="传递函数分母系数 (Denominator)",
|
|
value="1,6,11,6",
|
|
placeholder="例如: 1,2,1",
|
|
info="💡 分母阶数通常高于或等于分子阶数"
|
|
)
|
|
|
|
# 创建功能选项卡
|
|
with gr.Tabs() as tabs:
|
|
with gr.TabItem("⏱️ 时域分析 (Time Domain)", id=0):
|
|
time_domain_ui = create_time_domain_tab()
|
|
with gr.TabItem("📊 频域分析 (Frequency Domain)", id=1):
|
|
freq_domain_ui = create_frequency_domain_tab()
|
|
with gr.TabItem("🎯 根轨迹 (Root Locus)", id=2):
|
|
root_locus_ui = create_root_locus_tab()
|
|
# ===== 新增:算例演示 Tab(位于根轨迹与智能问答之间) =====
|
|
with gr.TabItem("🧪 算例演示 (Case Demo)", id=3):
|
|
case_demo_ui = create_case_demo_tab()
|
|
with gr.TabItem("🤖 智能问答 (Q&A)", id=4):
|
|
chatbot_ui = create_chatbot_tab()
|
|
|
|
# 2. 绑定事件逻辑
|
|
# --- 通用函数 ---
|
|
# 每次操作前更新用户活跃状态
|
|
def wrap_with_activity_update(fn, sid):
|
|
update_user_activity(sid)
|
|
# 使用 partial 将 session_id 绑定到函数上
|
|
# 这样Gradio调用时就不需要显式传递session_id了
|
|
return partial(fn, session_id=sid)
|
|
|
|
# --- 时域分析事件 ---
|
|
time_domain_ui["confirm_button"].click(
|
|
fn=display_transfer_function,
|
|
inputs=[num_input, den_input],
|
|
outputs=[time_domain_ui["tf_display"]]
|
|
).then(lambda: get_online_status_html(), outputs=online_counter)
|
|
|
|
time_domain_ui["analyze_button"].click(
|
|
fn=time_domain_analysis,
|
|
inputs=[num_input, den_input],
|
|
outputs=[time_domain_ui["output_plot"], time_domain_ui["output_metrics"]]
|
|
).then(lambda: get_online_status_html(), outputs=online_counter)
|
|
|
|
# --- 频域分析事件 ---
|
|
def update_frequency_analysis_wrapper(num, den, log_k):
|
|
k = 10**log_k
|
|
fig, metrics, tf_latex, stability = frequency_domain_analysis(num, den, k)
|
|
return fig, metrics, tf_latex, stability, k, get_online_status_html()
|
|
|
|
freq_inputs = [num_input, den_input, freq_domain_ui["log_k_slider"]]
|
|
freq_outputs = [
|
|
freq_domain_ui["plot_output"],
|
|
freq_domain_ui["metrics_display"],
|
|
freq_domain_ui["tf_display"],
|
|
freq_domain_ui["stability_display"],
|
|
freq_domain_ui["k_number_display"],
|
|
online_counter
|
|
]
|
|
freq_domain_ui["log_k_slider"].release(
|
|
fn=update_frequency_analysis_wrapper,
|
|
inputs=freq_inputs,
|
|
outputs=freq_outputs
|
|
)
|
|
|
|
# --- 根轨迹分析事件 ---
|
|
def update_rl_view_wrapper(log_k, num, den):
|
|
fig, poles, k_val = root_locus_analysis(num, den, log_k)
|
|
return fig, poles, k_val, get_online_status_html()
|
|
|
|
rl_inputs = [root_locus_ui["log_k_slider"], num_input, den_input]
|
|
rl_outputs = [
|
|
root_locus_ui["plot_output"],
|
|
root_locus_ui["poles_display"],
|
|
root_locus_ui["k_number_display"],
|
|
online_counter
|
|
]
|
|
root_locus_ui["log_k_slider"].release(
|
|
fn=update_rl_view_wrapper,
|
|
inputs=rl_inputs,
|
|
outputs=rl_outputs
|
|
)
|
|
|
|
# 当输入框变化时,也更新频域和根轨迹(如果它们是当前可见的)
|
|
def update_all_on_tf_change(num, den, log_k_freq, log_k_rl):
|
|
# 更新频域
|
|
k_freq = 10**log_k_freq
|
|
fig_freq, metrics, tf_latex, stability = frequency_domain_analysis(num, den, k_freq)
|
|
|
|
# 更新根轨迹
|
|
fig_rl, poles, k_val_rl = root_locus_analysis(num, den, log_k_rl)
|
|
|
|
return (
|
|
fig_freq, metrics, tf_latex, stability, k_freq,
|
|
fig_rl, poles, k_val_rl,
|
|
get_online_status_html()
|
|
)
|
|
|
|
tf_change_inputs = [num_input, den_input, freq_domain_ui["log_k_slider"], root_locus_ui["log_k_slider"]]
|
|
tf_change_outputs = freq_outputs[:-1] + rl_outputs[:-1] + [online_counter]
|
|
|
|
num_input.change(fn=update_all_on_tf_change, inputs=tf_change_inputs, outputs=tf_change_outputs)
|
|
den_input.change(fn=update_all_on_tf_change, inputs=tf_change_inputs, outputs=tf_change_outputs)
|
|
|
|
# ===== 新增:算例演示事件包装器 =====
|
|
def run_case_demo_wrapper(sim_time, dt, initial_soc, initial_engine_power, profile, rpm_scale, load_scale, sid):
|
|
update_user_activity(sid)
|
|
fig, summary, table_data = run_case_demo(
|
|
sim_time_s=sim_time,
|
|
dt=dt,
|
|
initial_soc_pct=initial_soc,
|
|
initial_engine_power_kw=initial_engine_power,
|
|
profile_name=profile,
|
|
rpm_scale=rpm_scale,
|
|
load_scale=load_scale
|
|
)
|
|
return fig, summary, table_data, get_online_status_html()
|
|
|
|
# ===== 新增:算例演示按钮事件绑定 =====
|
|
case_demo_ui["run_button"].click(
|
|
fn=run_case_demo_wrapper,
|
|
inputs=[
|
|
case_demo_ui["sim_time"],
|
|
case_demo_ui["dt"],
|
|
case_demo_ui["initial_soc"],
|
|
case_demo_ui["initial_engine_power"],
|
|
case_demo_ui["profile"],
|
|
case_demo_ui["rpm_scale"],
|
|
case_demo_ui["load_scale"],
|
|
session_id
|
|
],
|
|
outputs=[
|
|
case_demo_ui["plot"],
|
|
case_demo_ui["summary"],
|
|
case_demo_ui["table"],
|
|
online_counter
|
|
]
|
|
)
|
|
|
|
# --- 聊天机器人事件 ---
|
|
async def chat_wrapper(message, history, sid):
|
|
update_user_activity(sid)
|
|
# chat_with_ai 是一个生成器,Gradio可以直接处理
|
|
async for response in chat_with_ai(message, history):
|
|
yield response
|
|
|
|
chatbot_ui["send_button"].click(
|
|
fn=chat_wrapper,
|
|
inputs=[chatbot_ui["chat_input"], chatbot_ui["chatbot"], session_id],
|
|
outputs=chatbot_ui["chatbot"]
|
|
).then(lambda: ("", get_online_status_html()), outputs=[chatbot_ui["chat_input"], online_counter])
|
|
|
|
chatbot_ui["chat_input"].submit(
|
|
fn=chat_wrapper,
|
|
inputs=[chatbot_ui["chat_input"], chatbot_ui["chatbot"], session_id],
|
|
outputs=chatbot_ui["chatbot"]
|
|
).then(lambda: ("", get_online_status_html()), outputs=[chatbot_ui["chat_input"], online_counter])
|
|
|
|
def clear_chat_wrapper(sid):
|
|
update_user_activity(sid)
|
|
return [], get_online_status_html()
|
|
|
|
chatbot_ui["clear_button"].click(
|
|
fn=clear_chat_wrapper,
|
|
inputs=[session_id],
|
|
outputs=[chatbot_ui["chatbot"], online_counter]
|
|
)
|
|
|
|
# --- 页面加载和定时器事件 ---
|
|
def on_page_load(sid):
|
|
update_user_activity(sid)
|
|
return get_online_status_html()
|
|
|
|
demo.load(fn=on_page_load, inputs=[session_id], outputs=[online_counter])
|
|
|
|
gr.Timer(10).tick(fn=get_online_status_html, outputs=online_counter)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
demo.queue().launch(
|
|
server_name=config.SERVER_NAME,
|
|
server_port=config.SERVER_PORT,
|
|
share=config.SHARE
|
|
)
|