import pandas as pd from config import START_DATE, END_DATE, PRECISION_CONFIG from stock_utils import * from getStockDiv2 import analyze_stock_dividend_and_price from getStockParam import getStockParam from .data_source import get_tushare_pro import datetime pro = get_tushare_pro() # --- 数据输出配置 --- OUTPUT_TO_CSV = True # 开关:是否输出到CSV OUTPUT_TO_MYSQL = False # 开关:是否输出到MySQL (需要额外配置数据库连接) # --- 数据库输出函数 (伪代码) --- def save_to_mysql(dataframe, table_name): """ 将DataFrame保存到MySQL数据库的伪代码函数。 实际使用时需要配置数据库连接。 """ if not OUTPUT_TO_MYSQL: print("MySQL输出已禁用。") return print(f"正在将数据保存到MySQL表 {table_name}...") # --- 伪代码开始 --- # import pymysql # 或其他数据库连接库 # # connection = pymysql.connect(host='your_host', # user='your_user', # password='your_password', # database='your_database', # charset='utf8mb4') # try: # with connection.cursor() as cursor: # # 创建表 (如果不存在) # # create_table_sql = "..." # # cursor.execute(create_table_sql) # # # 遍历DataFrame并插入数据 # for index, row in dataframe.iterrows(): # # 注意:需要处理SQL注入风险,最好使用参数化查询 # insert_sql = f"INSERT INTO {table_name} (...) VALUES (...)" # cursor.execute(insert_sql, tuple(row)) # connection.commit() # print(f"成功保存 {len(dataframe)} 条记录到MySQL。") # except Exception as e: # print(f"保存到MySQL时出错: {e}") # connection.rollback() # finally: # connection.close() # --- 伪代码结束 --- print("MySQL保存操作完成 (伪代码)。") # --- CSV输出函数 --- def save_to_csv(dataframe, filename, mode='w', header=True): """将DataFrame保存到CSV文件""" if not OUTPUT_TO_CSV: print("CSV输出已禁用。") return try: dataframe.to_csv(filename, index=False, encoding='utf-8', mode=mode, header=header) print(f"成功保存 {len(dataframe)} 条记录到 {filename}") except Exception as e: print(f"保存CSV文件 {filename} 时出错: {e}") ''' 写一个函数,调用上面的函数,传入股票代码和起止日期,计算div_yield的最大值,最小值,均值,最新值,以及标准差,返回pandas的Series ''' def calculate_dividend_yield_stats(ts_code, start_date=START_DATE, end_date=END_DATE): """ 计算股息率的统计指标 参数: ts_code (str): 股票代码 start_date (str): 起始日期 end_date (str): 结束日期 返回: pd.Series: 包含股息率统计指标的Series """ df = analyze_stock_dividend_and_price(ts_code, start_date, end_date) if df.empty or 'div_yield' not in df.columns: return pd.Series(dtype=float) div_yield_data = df['div_yield'] stats = pd.Series({ 'ts_code': ts_code, 'max': div_yield_data.max(), 'min': div_yield_data.min(), 'mean': div_yield_data.mean(), 'latest': div_yield_data.iloc[-1] if len(div_yield_data) > 0 else 0, 'std': div_yield_data.std(), 'zero_count': (div_yield_data == 0).sum() }) return stats # --- 主处理函数 --- def process_stock_dividend_batch(start_dt='20200101', end_dt='20251231', dataInterval=10, exchange='SSE'): """ 批量处理股票股息率数据,并分批输出。 """ print(f"开始处理 {exchange} 交易所的股票数据,时间范围: {start_dt} - {end_dt}") try: df_stocks = get_stock_basic(exchange=exchange) if df_stocks.empty: print("未获取到股票基础数据") return except Exception as e: print(f"获取股票基础数据时出错: {e}") return results_batch = [] all_results = [] processed_count = 0 total_stocks = len(df_stocks) csv_filename = f"{exchange}_dividend_yield_stats_batch.csv" csv_first_write = True # 用于控制CSV文件头只写入一次 for idx, row in df_stocks.iterrows(): stock_code = row['ts_code'] print(f"正在处理 {stock_code} ({idx+1}/{total_stocks})") try: raw_stats = calculate_dividend_yield_stats(stock_code, start_date=start_dt, end_date=end_dt) # 确保我们最终得到一个 Series (统计数据) stats_series = None if isinstance(raw_stats, pd.Series): # 如果直接返回了 Series stats_series = raw_stats elif isinstance(raw_stats, pd.DataFrame) and not raw_stats.empty: # 如果返回了 DataFrame,则取第一行 stats_series = raw_stats.iloc[0] else: # 如果返回了空DataFrame, None, 或其他类型 print(f"股票 {stock_code} 的 calculate_dividend_yield_stats 返回了空数据或无效类型 ({type(raw_stats)})。跳过...") continue yesterday = (datetime.datetime.now() - datetime.timedelta(days=1)).strftime("%Y%m%d") stockParam = getStockParam(stock_code, START_DATE=yesterday, END_DATE=yesterday) # 选择需要的字段与下方合并 if not stockParam.empty: selected_params_dict = { 'ts_code': stockParam['ts_code'].iloc[0], 'trade_date': stockParam['trade_date'].iloc[0], 'total_mv': stockParam['total_mv'].iloc[0], 'circ_mv': stockParam['circ_mv'].iloc[0] } else: selected_params_dict = { 'ts_code': stock_code, 'trade_date': None, 'total_mv': None,'circ_mv': None} # --- 修改点2:简化合并逻辑 --- # 现在 stats_series 已确认是 Series,可以直接合并 # 使用字典合并确保结果清晰且 stats 覆盖同名项 combined_dict = {**row.to_dict(), **stats_series.to_dict(), **selected_params_dict} combined_series = pd.Series(combined_dict) print(f"合并结果: {combined_series}") print(f"{stock_code}-{start_dt} -- {end_dt} 处理完成。") results_batch.append(combined_series) all_results.append(combined_series) except Exception as e: print(f"处理股票 {stock_code} 时出错: {e}") # 可以选择在这里添加错误记录到 results_batch 或 all_results continue processed_count += 1 # 检查是否达到批次大小 if processed_count % dataInterval == 0 and results_batch: print(f"已处理 {processed_count} 支股票,达到批次大小 {dataInterval},开始输出...") df_batch = pd.DataFrame(results_batch) # 输出到CSV (追加模式) save_to_csv(df_batch, csv_filename, mode='a', header=csv_first_write) if csv_first_write: csv_first_write = False # 之后的批次不再写入header # 输出到MySQL (伪代码) save_to_mysql(df_batch, 'stock_dividend_stats') # 清空批次缓存 results_batch = [] print("--- 批次处理完成 ---") # 处理最后一批不足dataInterval的数据 if results_batch: print(f"处理剩余 {len(results_batch)} 支股票...") df_final_batch = pd.DataFrame(results_batch) save_to_csv(df_final_batch, csv_filename, mode='a', header=csv_first_write) save_to_mysql(df_final_batch, 'stock_dividend_stats') print("--- 最终批次处理完成 ---") # (可选) 将所有结果一次性保存到一个完整的CSV文件 if all_results: final_csv_filename = f"{exchange}_dividend_yield_stats_final.csv" print(f"正在保存所有 {len(all_results)} 条结果到 {final_csv_filename}...") df_final = pd.DataFrame(all_results) save_to_csv(df_final, final_csv_filename, mode='w', header=True) save_to_mysql(df_final, 'stock_dividend_stats_final') print("所有数据处理并保存完成。") else: print("无有效结果数据可保存") # --- 程序入口 --- if __name__ == "__main__": # 可以通过修改这里的参数来调用函数 process_stock_dividend_batch(start_dt='20150101', end_dt='20251231', dataInterval=10, exchange='SSE') process_stock_dividend_batch(start_dt='20150101', end_dt='20251231', dataInterval=10, exchange='SZSE')