Smart-Trader-EA commited on
Commit
c6ef3b3
·
1 Parent(s): 5e4ecc0

升级为本地数据分析版

Browse files
Files changed (2) hide show
  1. app.py +138 -42
  2. requirements.txt +0 -1
app.py CHANGED
@@ -1,54 +1,150 @@
1
  import gradio as gr
2
- import yfinance as yf
3
  import pandas as pd
4
  import plotly.graph_objects as go
5
  from prophet import Prophet
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6
 
7
  def analyze_stock(ticker):
8
- stock = yf.Ticker(ticker)
9
- hist = stock.history(period="1y")
10
- if hist.empty:
11
- return "股票代码无效,请检查后重试(示例:AAPL, TSLA, 600036.SS)", None, None
12
-
13
- fig = go.Figure(data=[go.Candlestick(x=hist.index,
14
- open=hist['Open'],
15
- high=hist['High'],
16
- low=hist['Low'],
17
- close=hist['Close'])])
18
- fig.update_layout(title=f"{ticker} 股票K线图", xaxis_title="日期", yaxis_title="价格")
19
-
20
- df = hist[['Close']].reset_index()
21
- df.columns = ['ds', 'y']
22
- model = Prophet(daily_seasonality=True, yearly_seasonality=True)
23
- model.fit(df)
24
- future = model.make_future_dataframe(periods=7)
25
- forecast = model.predict(future)
26
-
27
- fig2 = go.Figure()
28
- fig2.add_trace(go.Scatter(x=df['ds'], y=df['y'], mode='lines', name='历史价格'))
29
- fig2.add_trace(go.Scatter(x=forecast['ds'], y=forecast['yhat'], mode='lines', name='预测价格'))
30
- fig2.update_layout(title=f"{ticker} 7天价格预测", xaxis_title="日期", yaxis_title="价格")
31
-
32
- hist['MA20'] = hist['Close'].rolling(20).mean()
33
- current_price = hist['Close'].iloc[-1]
34
- ma20 = hist['MA20'].iloc[-1]
35
- signal = "📈 看涨" if current_price > ma20 else "📉 看跌"
36
-
37
- result_text = f"当前价格: ${current_price:.2f}\n20日均线: ${ma20:.2f}\n信号: {signal}"
38
- return result_text, fig, fig2
39
-
40
- with gr.Blocks(title="股票AI分析") as demo:
41
- gr.Markdown("# 📈 股票AI分析系统")
42
- gr.Markdown("输入股票代码(美股直接输入,A股加.SS后缀,如`600036.SS`)")
43
-
44
- ticker_input = gr.Textbox(label="股票代码", value="AAPL")
45
- analyze_btn = gr.Button("分析", variant="primary")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46
 
47
  signal_output = gr.Textbox(label="分析结果")
48
- kchart = gr.Plot(label="K线图")
49
- pred_chart = gr.Plot(label="价格预测")
 
50
 
51
- gr.Markdown("### 示例代码:\n- 美股: `AAPL`, `TSLA`, `GOOGL`\n- A股: `600036.SS`, `000001.SZ`\n- 港股: `00700.HK`")
 
 
 
 
52
 
53
  analyze_btn.click(
54
  fn=analyze_stock,
 
1
  import gradio as gr
 
2
  import pandas as pd
3
  import plotly.graph_objects as go
4
  from prophet import Prophet
5
+ import os
6
+
7
+
8
+ DATA_DIR = "data"
9
+ available_tickers = {}
10
+
11
+
12
+ if os.path.exists(DATA_DIR):
13
+ for filename in os.listdir(DATA_DIR):
14
+ if filename.endswith(".csv"):
15
+ ticker_name = filename.replace(".csv", "").upper()
16
+ file_path = os.path.join(DATA_DIR, filename)
17
+ try:
18
+ # 尝试不同编码读取
19
+ for encoding in ['utf-8', 'gbk', 'latin1']:
20
+ try:
21
+ df = pd.read_csv(file_path, encoding=encoding)
22
+ break
23
+ except:
24
+ continue
25
+
26
+ # 识别日期列
27
+ date_cols = [col for col in df.columns if 'date' in col.lower() or 'time' in col.lower()]
28
+ if date_cols:
29
+ df[date_cols[0]] = pd.to_datetime(df[date_cols[0]])
30
+ df.set_index(date_cols[0], inplace=True)
31
+
32
+ available_tickers[ticker_name] = {
33
+ "file": file_path,
34
+ "data": df
35
+ }
36
+ print(f"成功加载: {ticker_name}")
37
+ except Exception as e:
38
+ print(f"加载失败 {filename}: {str(e)}")
39
+ else:
40
+ print(f"警告: 数据目录不存在 - {DATA_DIR}")
41
 
42
  def analyze_stock(ticker):
43
+ if not any(available_tickers):
44
+ return "❌ 错误: 没有找到任何数据文件。请检查data/目录", None, None
45
+
46
+
47
+ ticker_upper = ticker.upper()
48
+ matched_ticker = None
49
+
50
+
51
+ if ticker_upper in available_tickers:
52
+ matched_ticker = ticker_upper
53
+ else:
54
+
55
+ for name in available_tickers.keys():
56
+ if ticker_upper in name or name in ticker_upper:
57
+ matched_ticker = name
58
+ break
59
+
60
+ if not matched_ticker:
61
+ return f"❌ 未找到匹配的数据: {ticker}\n可用数据: {', '.join(available_tickers.keys())}", None, None
62
+
63
+ try:
64
+ hist = available_tickers[matched_ticker]["data"]
65
+
66
+
67
+ required_cols = ['Open', 'High', 'Low', 'Close']
68
+ if not all(col in hist.columns for col in required_cols):
69
+ return f"❌ 数据格式错误: 缺少必要列。请确保CSV包含: {', '.join(required_cols)}", None, None
70
+
71
+
72
+ fig = go.Figure(data=[go.Candlestick(x=hist.index,
73
+ open=hist['Open'],
74
+ high=hist['High'],
75
+ low=hist['Low'],
76
+ close=hist['Close'])])
77
+ fig.update_layout(title=f"{matched_ticker} 股票K线图 (本地数据)", xaxis_title="日期", yaxis_title="价格")
78
+
79
+
80
+ df = hist[['Close']].reset_index()
81
+ df.columns = ['ds', 'y']
82
+ model = Prophet(daily_seasonality=True, yearly_seasonality=True)
83
+ model.fit(df)
84
+ future = model.make_future_dataframe(periods=30) # 30天预测
85
+ forecast = model.predict(future)
86
+
87
+ fig2 = go.Figure()
88
+ fig2.add_trace(go.Scatter(x=df['ds'], y=df['y'], mode='lines', name='历史价格'))
89
+ fig2.add_trace(go.Scatter(x=forecast['ds'], y=forecast['yhat'], mode='lines', name='预测价格'))
90
+ fig2.update_layout(title=f"{matched_ticker} 30天价格预测", xaxis_title="日期", yaxis_title="价格")
91
+
92
+
93
+ hist['MA20'] = hist['Close'].rolling(20).mean()
94
+ hist['MA50'] = hist['Close'].rolling(50).mean()
95
+ current_price = hist['Close'].iloc[-1]
96
+ ma20 = hist['MA20'].iloc[-1]
97
+ ma50 = hist['MA50'].iloc[-1]
98
+
99
+
100
+ if current_price > ma20 > ma50:
101
+ signal = "📈 强烈看涨 (黄金交叉)"
102
+ elif current_price > ma20:
103
+ signal = "📈 看涨"
104
+ elif current_price < ma20 < ma50:
105
+ signal = "📉 强烈看跌 (死亡交叉)"
106
+ else:
107
+ signal = "🔄 震荡"
108
+
109
+ result_text = (
110
+ f"📊 {matched_ticker} 分析结果\n"
111
+ f"💰 当前价格: ${current_price:.2f}\n"
112
+ f"📈 20日均线: ${ma20:.2f}\n"
113
+ f"📉 50日均线: ${ma50:.2f}\n"
114
+ f"🎯 信号: {signal}\n"
115
+ f"💾 数据来源: 本地文件 ({len(hist)} 条记录)"
116
+ )
117
+
118
+ return result_text, fig, fig2
119
+
120
+ except Exception as e:
121
+ return f"❌ 分析错误: {str(e)}", None, None
122
+
123
+ def list_available_data():
124
+ if not available_tickers:
125
+ return "暂无可用数据文件"
126
+ return "可用数据: " + ", ".join(available_tickers.keys())
127
+
128
+ with gr.Blocks(title="股票AI分析 (本地数据版)") as demo:
129
+ gr.Markdown("# 📈 股票AI分析系统 (本地历史数据版)")
130
+ gr.Markdown("✅ 优势: 无需网络,数据稳定,适合中长期分析")
131
+
132
+ data_status = gr.Textbox(label="数据状态", value=list_available_data(), interactive=False)
133
+
134
+ with gr.Row():
135
+ ticker_input = gr.Textbox(label="股票代码/名称", value="EURUSD")
136
+ analyze_btn = gr.Button("分析", variant="primary")
137
 
138
  signal_output = gr.Textbox(label="分析结果")
139
+ with gr.Row():
140
+ kchart = gr.Plot(label="K线图")
141
+ pred_chart = gr.Plot(label="价格预测")
142
 
143
+ gr.Markdown("### 使用指南:\n"
144
+ "1. 输入货币对名称(例如: EURUSD)\n"
145
+ "2. 系统自动从本地数据加载\n"
146
+ "3. 查看技术分析和30天预测\n"
147
+ "4. 定期更新data/目录中的CSV文件")
148
 
149
  analyze_btn.click(
150
  fn=analyze_stock,
requirements.txt CHANGED
@@ -1,5 +1,4 @@
1
  gradio==4.29.0
2
- yfinance==0.2.37
3
  pandas==2.2.2
4
  numpy==1.26.4
5
  plotly==5.22.0
 
1
  gradio==4.29.0
 
2
  pandas==2.2.2
3
  numpy==1.26.4
4
  plotly==5.22.0