Pandas单元格分组格式化优化:排序、分隔线与标记清除
Pandas DataFrame 格式化优化方案
原始DataFrame
import pandas as pd import numpy as np inp_df = pd.DataFrame( [ ["a1", "b1", "c1", "gbt", "auc", 82.5, 80.1, 83.6], ["a1", "b1", "c1", "gbt", "pr@5%", 0.3, 0.2, 0.4], ["a1", "b1", "c1", "gbt", "re@5%", 60.2, 58.1, 61.3], ["a1", "b1", "c1", "rnn", "auc", 84.1, 83.8, 84.5], ["a1", "b1", "c1", "rnn", "pr@5%", 0.5, 0.4, 0.6], ["a1", "b1", "c1", "rnn", "re@5%", 61.5, 61.4, 61.7], ["a1", "b1", "c1", "llm", "auc", 84.3, 84.1, 84.6], ["a1", "b1", "c1", "llm", "pr@5%", 0.8, 0.7, 0.9], ["a1", "b1", "c1", "llm", "re@5%", 61.2, 61.1, 61.3], ["a1", "b1", "c2", "gbt", "auc", 82.5, 80.1, 83.6], ["a1", "b1", "c2", "gbt", "pr@5%", 0.3, 0.2, 0.4], ["a1", "b1", "c2", "gbt", "re@5%", 60.2, 58.1, 61.3], ["a1", "b1", "c2", "llm", "auc", 84.3, 84.1, 84.6], ["a1", "b1", "c2", "llm", "pr@5%", 0.8, 0.7, 0.9], ["a1", "b1", "c2", "llm", "re@5%", 61.2, 61.1, 61.3], ], columns=["A","B","C","model","metric","val","val_lo","val_hi"] )
格式化需求
- 对每个
metric(如auc),将val值最高的model对应单元格设为粗体 - 高亮同一
(A,B,C)分组内置信区间(val_lo,val_hi)重叠的所有模型单元格 - 每组模型后添加分隔线
已实现的部分代码
数据预处理
cols = ["val","val_lo","val_hi"] inp_df["value"] = list(inp_df[cols].to_records(index=False)) inp_df.drop(columns=cols, inplace=True) out_df = inp_df.pivot(index=inp_df.columns[:4], columns="metric", values="value")\ .reset_index().rename_axis(None, axis=1)
初步格式化逻辑
mets = ["auc","pr@5%","re@5%"] def flag(block): out = [block["model"].values.tolist()] for met in mets: val,lo,hi = map(np.array, zip(*block[met].values)) maxind = val.argmax() overlapbool = np.logical_and(hi[maxind]>=lo, lo[maxind]<=hi) overlapinds = set(np.where(overlapbool)[0]) if overlapbool.sum()>1 else set() curr = list() for n,(x,y,z) in enumerate(zip(val,lo,hi)): cell = f"{x:.1f} ({y:.1f}-{z:.1f})" if n==maxind: cell += "*" if n in overlapinds: cell += "†" curr.append(cell) out.append(curr) return pd.DataFrame(zip(*out), columns=["model"]+mets) styled_df = out_df.groupby(["A","B","C"]).apply(flag).droplevel(-1).set_index(["model"], append=True)\ .style.applymap(lambda val:"font-weight:bold" if "*" in val else None)\ .applymap(lambda val:f"background-color:beige" if "†" in val else None)
待解决的问题
- 模型无法按
[gbt, rnn, llm]指定顺序排列 - 无法在每组模型后添加分隔线(如
(a1,b1,c1)与(a1,b1,c2)之间) - 需要去除用于标记格式化的
*和†字符
完整优化解决方案
import pandas as pd import numpy as np # 1. 设置模型有序分类,确保排序生效 model_order = ["gbt", "rnn", "llm"] inp_df["model"] = pd.Categorical(inp_df["model"], categories=model_order, ordered=True) # 数据预处理 cols = ["val","val_lo","val_hi"] inp_df["value"] = list(inp_df[cols].to_records(index=False)) inp_df.drop(columns=cols, inplace=True) out_df = inp_df.pivot(index=inp_df.columns[:4], columns="metric", values="value")\ .reset_index().rename_axis(None, axis=1) # 按模型顺序排序 out_df = out_df.sort_values("model") mets = ["auc","pr@5%","re@5%"] def flag(block): # 按指定模型顺序排序分组内数据 block = block.sort_values("model") out_rows = [] models = block["model"].tolist() # 先收集每个metric的样式标记和单元格内容 metric_data = {} for met in mets: val, lo, hi = map(np.array, zip(*block[met].values)) maxind = val.argmax() overlapbool = np.logical_and(hi[maxind] >= lo, lo[maxind] <= hi) overlapinds = set(np.where(overlapbool)[0]) if overlapbool.sum() > 1 else set() cells = [] is_bold = [] is_highlight = [] for n, (x, y, z) in enumerate(zip(val, lo, hi)): cells.append(f"{x:.1f} ({y:.1f}-{z:.1f})") is_bold.append(n == maxind) is_highlight.append(n in overlapinds) metric_data[met] = {"cells": cells, "is_bold": is_bold, "is_highlight": is_highlight} # 组装每行数据和样式标记 for idx in range(len(models)): row = {"model": models[idx]} row_styles = {"model": "", **{met: "" for met in mets}} for met in mets: row[met] = metric_data[met]["cells"][idx] if metric_data[met]["is_bold"][idx]: row_styles[met] += "font-weight:bold;" if metric_data[met]["is_highlight"][idx]: row_styles[met] += "background-color:beige;" out_rows.append((row, row_styles)) # 返回数据框和样式信息 df = pd.DataFrame([r[0] for r in out_rows], columns=["model"] + mets) df._styles = [r[1] for r in out_rows] return df # 分组处理 grouped_result = out_df.groupby(["A","B","C"], group_keys=True).apply(flag) final_df = grouped_result.droplevel(-1).set_index(["model"], append=True) # 2. 添加分组分隔线 def add_group_separator(styler): # 获取每个分组的最后一行索引 group_end_indices = [] for name, group in final_df.groupby(level=[0,1,2]): group_end_indices.append(group.index[-1]) # 设置表格样式,给分组最后一行添加底部边框 styler.set_table_styles([ {'selector': f'tbody tr:nth-child({i+1})', 'props': [('border-bottom', '2px solid #000')]} for i, idx in enumerate(final_df.index) if idx in group_end_indices ]) return styler # 3. 应用样式,不需要依赖标记字符 def apply_styles(styler): # 从自定义的_styles属性中读取样式 styles = [] for idx in final_df.index: row_idx = final_df.index.get_loc(idx) group_row_idx = row_idx % len(grouped_result.iloc[row_idx//len(model_order)]) row_style = grouped_result.iloc[row_idx//len(model_order)]._styles[group_row_idx] row_styles = [] for col in final_df.columns: row_styles.append(row_style[col]) styles.append(row_styles) styler.apply(lambda x: styles, axis=None) return styler # 生成最终样式表格 final_styled = final_df.style.pipe(add_group_separator).pipe(apply_styles) final_styled
关键优化点说明
- 模型排序:将
model列设为pd.Categorical并指定有序分类,后续排序和分组时都会遵循该顺序;分组内也额外做了排序确保顺序正确。 - 分组分隔线:通过
set_table_styles定位每个分组的最后一行,添加底部粗边框实现分隔效果。 - 去除标记字符:不再用
*和†标记样式,而是在flag函数中单独记录每个单元格的样式状态,直接应用到表格样式中,避免残留标记字符。
内容的提问来源于stack exchange,提问作者ironv
相关产品推荐
相关产品推荐

