""" 候选股过滤与数据加载 """ import pandas as pd from datetime import date, timedelta from core.scoring.models import ( KlineStock, StockInfo, IndustryMapping, KlineIndex, MarketRegimeDaily, SectorFeaturesDaily, ) from core.scoring.config import ( FILTER_MIN_CLOSE, FILTER_MAX_CLOSE, REQUIRE_DAYS, WINDOW_180D, HS300_CODE, ) from core.logger import LogLevel, PrintLog class DataContext: """特征计算所需的全部数据上下文""" def __init__(self, trade_date: date): self.trade_date = trade_date self.require_days = REQUIRE_DAYS # 原始数据 self.kline: pd.DataFrame = None # 候选股 K线 (180d) self.stock_info: dict = {} # code → StockInfo dict self.industry_map: dict = {} # code → industry_name self.hs300_kline: pd.DataFrame = None # HS300 K线 (180d) self.market_regime: pd.DataFrame = None # 市场状态 (180d) self.sector_features: pd.DataFrame = None # 行业指数 (180d) # 候选股列表 self.candidates: list[str] = [] self.excluded: dict[str, str] = {} # code → reason def _get_stock_codes_for_date(trade_date: date) -> list: """获取评分日所有符合条件的股票代码(非ST/退市)""" rows = (StockInfo .select(StockInfo.code, StockInfo.listing_status) .where(StockInfo.listing_status.not_in(['delisted', 'ST'])) .dicts()) return [row['code'] for row in rows] def _check_kline_sufficiency(stock_code: str, trade_date: date) -> tuple[bool, str, float]: """检查单只股票的K线数据是否满足评分条件""" # 查询最近 REQUIRE_DAYS + 30 (留缓冲) 个交易日 start = trade_date - timedelta(days=(REQUIRE_DAYS + 60) * 2) rows = (KlineStock .select(KlineStock.trade_date, KlineStock.close) .where( (KlineStock.stock_code == stock_code) & (KlineStock.trade_date >= start) & (KlineStock.trade_date <= trade_date) ) .order_by(KlineStock.trade_date.desc()) .dicts()) if not rows: return False, '无K线数据', 0 # 过滤有效收盘价 (close > 0) valid_rows = [r for r in rows if r['close'] and r['close'] > 0] if len(valid_rows) < REQUIRE_DAYS: return False, f'交易天数不足({len(valid_rows)}<{REQUIRE_DAYS})', 0 latest_close = valid_rows[0]['close'] # 价格区间检查 if latest_close < FILTER_MIN_CLOSE: return False, f'价格过低({latest_close:.2f}<{FILTER_MIN_CLOSE})', latest_close if latest_close > FILTER_MAX_CLOSE: return False, f'价格过高({latest_close:.2f}>{FILTER_MAX_CLOSE})', latest_close return True, '', latest_close def load_candidates(trade_date: date) -> DataContext: """ 加载评分日候选股及全部所需数据。 返回 DataContext,包含候选股列表和预加载的原始数据。 """ ctx = DataContext(trade_date) start_180 = trade_date - timedelta(days=365) # 取约1年数据覆盖180个交易日 # 1. 获取非ST/退市的全部股票 all_codes = [r.split('.')[0] for r in _get_stock_codes_for_date(trade_date)] PrintLog(LogLevel.INFO, f'[validator] 全市场有效股票: {len(all_codes)} 只') # 2. 批量加载 KlineStock (180d 窗口) raw_codes = all_codes # 使用纯数字代码查询 kline_rows = (KlineStock .select() .where( (KlineStock.stock_code.in_(raw_codes)) & (KlineStock.trade_date >= start_180) & (KlineStock.trade_date <= trade_date) ) .order_by(KlineStock.stock_code, KlineStock.trade_date) .dicts()) kline_df = pd.DataFrame(kline_rows) if kline_df.empty: PrintLog(LogLevel.WARNING, '[validator] KlineStock 无数据') return ctx kline_df['trade_date'] = pd.to_datetime(kline_df['trade_date']) PrintLog(LogLevel.INFO, f'[validator] K线原始数据: {len(kline_df)} 行') # 3. 逐股过滤 candidates = [] excluded = {} grouped = kline_df.groupby('stock_code') for code, group in grouped: g_sorted = group.sort_values('trade_date', ascending=False) valid_rows = g_sorted[g_sorted['close'].notna() & (g_sorted['close'] > 0)] if len(valid_rows) < REQUIRE_DAYS: excluded[code] = f'交易天数不足({len(valid_rows)}<{REQUIRE_DAYS})' continue latest_close = valid_rows.iloc[0]['close'] if latest_close < FILTER_MIN_CLOSE: excluded[code] = f'价格过低({latest_close:.2f}<{FILTER_MIN_CLOSE})' continue if latest_close > FILTER_MAX_CLOSE: excluded[code] = f'价格过高({latest_close:.2f}>{FILTER_MAX_CLOSE})' continue candidates.append(code) PrintLog(LogLevel.INFO, f'[validator] 候选: {len(candidates)} 通过, {len(excluded)} 排除') # 4. 裁剪K线到只含候选股 (保留最近180日) kline_df = kline_df[kline_df['stock_code'].isin(candidates)].copy() cut_date = trade_date - timedelta(days=365) kline_df = kline_df[kline_df['trade_date'] >= pd.Timestamp(cut_date)] ctx.kline = kline_df ctx.candidates = candidates ctx.excluded = excluded # 5. 加载 StockInfo stock_rows = (StockInfo .select() .where(StockInfo.code.in_([f'{c}.SH' for c in candidates] + [f'{c}.SZ' for c in candidates] + [f'{c}.BJ' for c in candidates])) .dicts()) ctx.stock_info = {r['code']: r for r in stock_rows} # 6. 加载行业映射: code → industry_name ind_rows = (IndustryMapping .select() .where(IndustryMapping.code.in_(candidates)) .dicts()) ctx.industry_map = {r['code']: r['industry_name'] for r in ind_rows} PrintLog(LogLevel.INFO, f'[validator] 行业映射: {len(ctx.industry_map)} 只') # 7. 加载 HS300 K线 hs300_rows = (KlineIndex .select() .where( (KlineIndex.index_code == HS300_CODE) & (KlineIndex.trade_date >= start_180) & (KlineIndex.trade_date <= trade_date) ) .order_by(KlineIndex.trade_date) .dicts()) ctx.hs300_kline = pd.DataFrame(hs300_rows) if not ctx.hs300_kline.empty: ctx.hs300_kline['trade_date'] = pd.to_datetime(ctx.hs300_kline['trade_date']) # 8. 加载市场状态 (180d) mkt_rows = (MarketRegimeDaily .select() .where( (MarketRegimeDaily.trade_date >= start_180) & (MarketRegimeDaily.trade_date <= trade_date) ) .order_by(MarketRegimeDaily.trade_date) .dicts()) ctx.market_regime = pd.DataFrame(mkt_rows) if not ctx.market_regime.empty: ctx.market_regime['trade_date'] = pd.to_datetime(ctx.market_regime['trade_date']) # 9. 加载行业指数 (180d) sec_rows = (SectorFeaturesDaily .select() .where( (SectorFeaturesDaily.trade_date >= start_180) & (SectorFeaturesDaily.trade_date <= trade_date) ) .order_by(SectorFeaturesDaily.trade_date) .dicts()) ctx.sector_features = pd.DataFrame(sec_rows) if not ctx.sector_features.empty: ctx.sector_features['trade_date'] = pd.to_datetime(ctx.sector_features['trade_date']) return ctx