You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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)

待解决的问题

  1. 模型无法按[gbt, rnn, llm]指定顺序排列
  2. 无法在每组模型后添加分隔线(如(a1,b1,c1)与(a1,b1,c2)之间)
  3. 需要去除用于标记格式化的*和†字符

完整优化解决方案

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

关键优化点说明

  1. 模型排序:将model列设为pd.Categorical并指定有序分类,后续排序和分组时都会遵循该顺序;分组内也额外做了排序确保顺序正确。
  2. 分组分隔线:通过set_table_styles定位每个分组的最后一行,添加底部粗边框实现分隔效果。
  3. 去除标记字符:不再用*和†标记样式,而是在flag函数中单独记录每个单元格的样式状态,直接应用到表格样式中,避免残留标记字符。

内容的提问来源于stack exchange,提问作者ironv

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.25 01:07:08