# 稀疏PCA特征提取与载荷分析
from sklearn.decomposition import SparsePCA
import seaborn as sns
# ── 时间段1:2018-2020 ──
X1 = window_pre.values.T # 转置为 n_samples x n_features 格式
spca_pre = SparsePCA(n_components=6, alpha=0.05, ridge_alpha=0.01, max_iter=800, random_state=42)
spca_pre.fit(X1)
# 提取变换后的主成分得分
sparse_scores_pre = spca_pre.transform(X1)
# 计算各成分解释方差
explained_vars_pre = np.var(sparse_scores_pre, axis=0)
sparse_ratio_pre = explained_vars_pre / explained_vars_pre.sum()
# ── 时间段2:2020-2021 ──
window_post = returns_log.loc['2020-02-01':'2021-12-31']
# 移除缺失较多的列
window_post.drop(columns=['600188.SH', '601919.SH'], inplace=True, errors='ignore')
window_post = window_post.astype('float64')
scaler_post = StandardScaler()
X2 = scaler_post.fit_transform(window_post).T # 转置后标准化
spca_post = SparsePCA(n_components=6, alpha=0.05, ridge_alpha=0.01, max_iter=800, random_state=42)
spca_post.fit(X2)
sparse_scores_post = spca_post.transform(X2)
explained_vars_post = np.var(sparse_scores_post, axis=0)
sparse_ratio_post = explained_vars_post / explained_vars_post.sum()
# ── 对比柱状图 ──
fig, axes = plt.subplots(1, 2, figsize=(14, 5))
labels = [f'成分{i+1}' for i in range(6)]
axes[0].bar(labels, sparse_ratio_pre, color='coral', edgecolor='black')
axes[0].set_title('稀疏PCA解释方差比例 (2018-2020)')
axes[0].set_ylabel('比例')
axes[0].set_ylim(0, 0.4)
axes[1].bar(labels, sparse_ratio_post, color='teal', edgecolor='black')
axes[1].set_title('稀疏PCA解释方差比例 (2020-2021)')
axes[1].set_ylim(0, 0.4)
plt.tight_layout()
plt.show()
# ── 载荷矩阵热力图 ──
loadings = spca_pre.components_ # shape: (6, n_features)
plt.figure(figsize=(12, 5))
sns.heatmap(
loadings,
cmap='Greys',
cbar_kws={'label': '载荷系数'},
xticklabels=5,
yticklabels=[f'成分_{i}' for i in range(6)],
linewidths=0.3
)
plt.title('Sparse PCA 载荷矩阵 (\u03bb=0.05)', fontsize=14)
plt.xlabel('特征索引', fontsize=12)
plt.ylabel('主成分', fontsize=12)
plt.tight_layout()
plt.show()
# 输出主成分1的载荷详情
pc1_loadings = pd.Series(loadings[0], name='成分1载荷')
top_features = pc1_loadings.abs().sort_values(ascending=False).head(15)
print("成分1前15大特征载荷:")
print(top_features)