Lou / app.py
Lousama's picture
Update app.py
5005d6f verified
Raw
History Blame Contribute Delete
8.73 kB
import os, tempfile, pathlib
import gradio as gr
import pandas as pd
import zipfile, time
from core_numeric import (
SUFFIX, CATEGORY_TITLES,
classify_row, _is_blank,
is_category_header_row, is_allowed_category_name,
find_header_row_by_probe
)
# ====== .xls 专用依赖(只在 .xls 分支用) ======
import xlrd # 读 .xls
from xlutils.copy import copy as xls_copy # 将 xlrd.Book -> xlwt.Workbook
import xlwt # 写回 .xls(BIFF8)
# ====== .xlsx 读写(现代路径) ======
import openpyxl
def _safe_basename(path: str) -> str:
return os.path.splitext(os.path.basename(path))[0]
def process_xls(in_path: str) -> str:
"""
在原工作表末尾新增 4 列:
IsTrueNumericMismatch / SrcNumbers / TgtNumbers / Note
并写回为 *_numeric_tagged.xls
"""
# 读取并保留样式信息
# 注意:formatting_info=True 仅对 .xls 生效,不能用于 .xlsx。:contentReference[oaicite:5]{index=5}
book = xlrd.open_workbook(in_path, formatting_info=True)
wbook = xls_copy(book)
sh_in = book.sheet_by_index(0)
sh_out = wbook.get_sheet(0)
nrows, ncols = sh_in.nrows, sh_in.ncols
# 找“Source/Target”表头行
def _cell(r, c):
try:
return sh_in.cell_value(r, c)
except Exception:
return ""
hdr = find_header_row_by_probe(_cell, max_scan=min(80, nrows))
# 追加 4 列写结果
base_col = ncols
headers = ["IsTrueNumericMismatch", "SrcNumbers", "TgtNumbers", "Note"]
if hdr >= 0:
for i, h in enumerate(headers):
sh_out.write(hdr, base_col + i, h)
# 遍历行并写入结果(保留原结构;不改变/删除任何行)
in_allowed = False
for r in range(hdr + 1 if hdr >= 0 else 0, nrows):
a = _cell(r, 0) # A: Category or "FileName (SegmentID)"
c = _cell(r, 2) # C: Source
d = _cell(r, 3) # D: Target
# 识别“类别标题行”
if is_category_header_row(a, c, d):
in_allowed = is_allowed_category_name(a)
continue
# 非允许类别(不是 Numeric mismatch)→ 不写结果
if not in_allowed:
continue
# 忽略纯空白行
if _is_blank(c) and _is_blank(d):
continue
# 计算并写入
tflag, src_s, tgt_s, note = classify_row(str(c or ""), str(d or ""))
sh_out.write(r, base_col + 0, "TRUE" if tflag else "FALSE")
sh_out.write(r, base_col + 1, src_s)
sh_out.write(r, base_col + 2, tgt_s)
sh_out.write(r, base_col + 3, note)
# 保存为 *_numeric_tagged.xls
out_path = os.path.join(
tempfile.gettempdir(),
f"{_safe_basename(in_path)}{SUFFIX}.xls"
)
wbook.save(out_path)
return out_path
def _read_table_generic(path: str) -> pd.DataFrame:
ext = pathlib.Path(path).suffix.lower()
if ext == ".csv":
return pd.read_csv(path, dtype=str, encoding="utf-8", keep_default_na=False)
elif ext == ".xls":
# pandas 读 .xls 也会用 xlrd 引擎;此路径仅用于 .xlsx 逻辑,.xls 我们已经直接“写回”。:contentReference[oaicite:6]{index=6}
return pd.read_excel(path, dtype=str, engine="xlrd")
else:
return pd.read_excel(path, dtype=str, engine="openpyxl")
def _reform_and_tag_df(df_raw: pd.DataFrame) -> pd.DataFrame:
"""面向 .xlsx/.csv:保留 'Numeric mismatch' 相关行,生成结构化输出表"""
# 探测表头行
def _probe(r, c):
try:
v = df_raw.iat[r, c]
return str(v) if pd.notna(v) else ""
except Exception:
return ""
hdr = find_header_row_by_probe(_probe, max_scan=min(80, len(df_raw)))
df = df_raw.iloc[hdr+1:].copy() if hdr >= 0 else df_raw.copy()
while df.shape[1] < 5:
df[df.shape[1]] = ""
colA, colB, colC, colD, colE = df.columns[0], df.columns[1], df.columns[2], df.columns[3], df.columns[4]
# 仅处理 Numeric mismatch 类别
keep, in_allowed = [], False
for _, row in df.iterrows():
a, c, d = row.get(colA, ""), row.get(colC, None), row.get(colD, None)
if is_category_header_row(a, c, d):
in_allowed = is_allowed_category_name(a); keep.append(False)
else:
keep.append(in_allowed)
df = df[pd.Series(keep, index=df.index)]
def is_data_row(row):
a_norm = str(row.get(colA, "")).strip().lower()
if a_norm in CATEGORY_TITLES: return False
return not (_is_blank(row.get(colC, None)) and _is_blank(row.get(colD, None)))
df = df[df.apply(is_data_row, axis=1)].copy()
# 计算 4 列结果
out_cols = ["IsTrueNumericMismatch", "SrcNumbers", "TgtNumbers", "Note"]
df_res = pd.DataFrame(columns=out_cols, index=df.index)
for i, row in df.iterrows():
tflag, src_s, tgt_s, note = classify_row(str(row.get(colC, "") or ""), str(row.get(colD, "") or ""))
df_res.at[i, "IsTrueNumericMismatch"] = bool(tflag)
df_res.at[i, "SrcNumbers"] = src_s
df_res.at[i, "TgtNumbers"] = tgt_s
df_res.at[i, "Note"] = note
# 拼出导出表(保留关注列 + 4 列结论)
out = pd.DataFrame({
"A": df[colA],
"Percent": df[colB],
"Source": df[colC],
"Target": df[colD],
"Comments": df[colE],
}).join(df_res)
return out
def process_xlsx_or_csv(in_path: str) -> str:
df_raw = _read_table_generic(in_path)
out_df = _reform_and_tag_df(df_raw)
out_path = os.path.join(
tempfile.gettempdir(),
f"{_safe_basename(in_path)}{SUFFIX}.xlsx"
)
out_df.to_excel(out_path, index=False)
# 设置筛选/冻结等属性(打开时生效)
wb = openpyxl.load_workbook(out_path); ws = wb.active
from openpyxl.utils import get_column_letter
ws.auto_filter.ref = f"A1:{get_column_letter(ws.max_column)}{ws.max_row}"
ws.freeze_panes = "A2"
wb.save(out_path)
return out_path
def process_files(files, pack_zip):
outputs = []
for f in files:
path = f.name
ext = pathlib.Path(path).suffix.lower()
if ext == ".xls":
outputs.append(process_xls(path))
elif ext in [".xlsx", ".csv"]:
outputs.append(process_xlsx_or_csv(path))
else:
raise gr.Error(f"不支持的文件类型:{ext}")
zip_path = None
if pack_zip and outputs:
# 以时间戳避免同名覆盖
stamp = time.strftime("%Y%m%d-%H%M%S")
zip_path = os.path.join(
tempfile.gettempdir(), f"XbenchAssistant_{stamp}.zip"
)
with zipfile.ZipFile(zip_path, "w", compression=zipfile.ZIP_DEFLATED) as zf:
for p in outputs:
zf.write(p, arcname=os.path.basename(p))
return outputs, (zip_path or "")
with gr.Blocks(title="XbenchAssistant") as demo:
gr.Image(value="logo.png", show_label=False, height=80)
gr.Markdown("### XbenchAssistant|Xbench QA 报告 · 数字一致性去噪与打标\n"
"上传 `.xls/.xlsx/.csv`;\n"
"- `.xls`:在原工作表**追加 4 列**后,导出为 `*_numeric_tagged.xls`;\n"
"- `.xlsx/.csv`:产出结构化 `*_numeric_tagged.xlsx`。")
in_files = gr.File(label="上传 Xbench 报告(可多选)", file_count="multiple",
file_types=[".xls", ".xlsx", ".csv"])
with gr.Row():
pack_zip = gr.Checkbox(value=True, label="打包为 ZIP(推荐)")
btn = gr.Button("开始处理", variant="primary")
out_files = gr.File(label="逐个下载(可多选点击)", interactive=False)
zip_dl = gr.DownloadButton(label="⬇️ 一键下载 ZIP(处理后打包)", visible=True)
# 现在返回两个输出:1) 文件列表 2) zip 路径
btn.click(process_files, inputs=[in_files, pack_zip], outputs=[out_files, zip_dl])
if __name__ == "__main__":
import os
# 从 HF 的 Secrets 里读取(Settings → Secrets,新建 APP_USER_1/APP_PASS_1/APP_USER_2/APP_PASS_2)
USER1 = os.getenv("APP_USER_1", "")
PASS1 = os.getenv("APP_PASS_1", "")
USER2 = os.getenv("APP_USER_2", "")
PASS2 = os.getenv("APP_PASS_2", "")
AUTH_PAIRS = []
if USER1 and PASS1:
AUTH_PAIRS.append((USER1, PASS1))
if USER2 and PASS2:
AUTH_PAIRS.append((USER2, PASS2))
# 登录页提示可自定义
AUTH_MSG = "ECI 内部工具|请输入用户名与密码"
# 多账户认证:Gradio 的 launch(auth=...) 支持元组或列表形式
demo.launch(
auth=AUTH_PAIRS if AUTH_PAIRS else None,
auth_message=AUTH_MSG,
ssr_mode=False,
show_error=True
)