Spaces:
Sleeping
Sleeping
khs commited on
Commit ·
d406944
1
Parent(s): b2d4551
Add kn-style calibration workflow
Browse files- README.md +46 -1
- app.py +167 -52
- calibration/model.json +7 -0
- scripts/train_calibration.py +62 -0
README.md
CHANGED
|
@@ -12,4 +12,49 @@ license: mit
|
|
| 12 |
short_description: 'an AIGC_detector by model: yuchuantian/AIGC_detector_zhv2'
|
| 13 |
---
|
| 14 |
|
| 15 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
short_description: 'an AIGC_detector by model: yuchuantian/AIGC_detector_zhv2'
|
| 13 |
---
|
| 14 |
|
| 15 |
+
# AIGC Detector (论文AIGC风险 + 知网对齐预测)
|
| 16 |
+
|
| 17 |
+
这是一个可直接部署到 Hugging Face Spaces 的 Gradio 项目。
|
| 18 |
+
|
| 19 |
+
## 功能
|
| 20 |
+
|
| 21 |
+
- 段落级 AIGC 风险分析
|
| 22 |
+
- 综合 AI 风险率
|
| 23 |
+
- 预测知网 AIGC 率(校准模式)
|
| 24 |
+
- 无校准文件时自动回退到原始风险率模式
|
| 25 |
+
|
| 26 |
+
## Space 直接部署
|
| 27 |
+
|
| 28 |
+
1. 上传本仓库到 Hugging Face Space(SDK 选 Gradio)。
|
| 29 |
+
2. Space 会自动读取 `requirements.txt` 安装依赖。
|
| 30 |
+
3. 默认即可运行,不需要额外环境变量。
|
| 31 |
+
|
| 32 |
+
## 知网对齐校准(推荐)
|
| 33 |
+
|
| 34 |
+
`calibration/model.json` 是校准模型文件。
|
| 35 |
+
|
| 36 |
+
- 默认提供的是透传占位模型(预测率=综合风险率)。
|
| 37 |
+
- 你可以用真实数据重新训练并覆盖它。
|
| 38 |
+
|
| 39 |
+
### 训练数据格式(CSV)
|
| 40 |
+
|
| 41 |
+
需要列:
|
| 42 |
+
|
| 43 |
+
- `overall`
|
| 44 |
+
- `p90`
|
| 45 |
+
- `high_ratio`
|
| 46 |
+
- `mid_ratio`
|
| 47 |
+
- `std`
|
| 48 |
+
- `kn_rate`(知网AIGC率,0~1)
|
| 49 |
+
|
| 50 |
+
### 训练命令
|
| 51 |
+
|
| 52 |
+
```bash
|
| 53 |
+
python3 scripts/train_calibration.py --input your_dataset.csv --target kn_rate --out calibration/model.json
|
| 54 |
+
```
|
| 55 |
+
|
| 56 |
+
训练完成后,把新的 `calibration/model.json` 提交到 Space,即可自动启用“知网对齐预测率”。
|
| 57 |
+
|
| 58 |
+
## 说明
|
| 59 |
+
|
| 60 |
+
本项目输出用于研究与预筛查,不代表任何官方系统结论。
|
app.py
CHANGED
|
@@ -1,104 +1,203 @@
|
|
| 1 |
-
import
|
| 2 |
-
import
|
| 3 |
-
import torch
|
| 4 |
-
import numpy as np
|
| 5 |
import re
|
|
|
|
|
|
|
| 6 |
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
|
|
|
| 11 |
|
| 12 |
MODEL_NAME = "yuchuantian/AIGC_detector_zhv2"
|
| 13 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 14 |
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
|
| 15 |
model = AutoModelForSequenceClassification.from_pretrained(MODEL_NAME)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
|
| 17 |
-
|
| 18 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
if len(p.strip()) > 50
|
| 24 |
-
]
|
| 25 |
|
| 26 |
-
|
| 27 |
-
|
|
|
|
| 28 |
|
| 29 |
-
|
| 30 |
-
|
| 31 |
|
| 32 |
-
|
| 33 |
|
| 34 |
-
return 1 - unique_ratio
|
| 35 |
|
| 36 |
-
def
|
| 37 |
-
|
|
|
|
|
|
|
| 38 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 39 |
lens = [len(s.strip()) for s in sents if len(s.strip()) > 0]
|
| 40 |
|
| 41 |
if len(lens) < 2:
|
| 42 |
-
return 0
|
|
|
|
|
|
|
| 43 |
|
| 44 |
-
return np.var(lens) / 1000
|
| 45 |
|
| 46 |
-
def detector_score(text):
|
| 47 |
inputs = tokenizer(
|
| 48 |
text,
|
| 49 |
-
return_tensors="pt",
|
| 50 |
truncation=True,
|
| 51 |
-
max_length=
|
|
|
|
|
|
|
|
|
|
|
|
|
| 52 |
)
|
| 53 |
|
| 54 |
with torch.no_grad():
|
| 55 |
outputs = model(**inputs)
|
| 56 |
|
| 57 |
probs = torch.softmax(outputs.logits, dim=-1)
|
|
|
|
| 58 |
|
| 59 |
-
return
|
| 60 |
|
| 61 |
-
def analyze_paragraph(text):
|
| 62 |
-
detector = detector_score(text)
|
| 63 |
|
|
|
|
|
|
|
| 64 |
repetition = calc_repetition(text)
|
| 65 |
-
|
| 66 |
variance = calc_sentence_variance(text)
|
| 67 |
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
+ (1 - min(variance, 1)) * 0.15
|
| 72 |
-
)
|
| 73 |
|
| 74 |
return {
|
| 75 |
"detector": detector,
|
| 76 |
"repetition": repetition,
|
| 77 |
"variance": variance,
|
| 78 |
-
"risk": risk
|
| 79 |
}
|
| 80 |
|
| 81 |
-
def analyze_pdf(pdf_file):
|
| 82 |
-
doc = fitz.open(pdf_file.name)
|
| 83 |
|
| 84 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 85 |
|
| 86 |
-
for
|
| 87 |
-
|
| 88 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 89 |
paragraphs = split_paragraphs(text)
|
|
|
|
|
|
|
| 90 |
|
| 91 |
-
|
|
|
|
| 92 |
|
|
|
|
| 93 |
risks = []
|
| 94 |
|
| 95 |
-
for p in paragraphs
|
| 96 |
score = analyze_paragraph(p)
|
| 97 |
-
|
| 98 |
risks.append(score["risk"])
|
| 99 |
|
| 100 |
level = "🟢"
|
| 101 |
-
|
| 102 |
if score["risk"] > 0.75:
|
| 103 |
level = "🔴"
|
| 104 |
elif score["risk"] > 0.55:
|
|
@@ -110,29 +209,45 @@ def analyze_pdf(pdf_file):
|
|
| 110 |
|
| 111 |
Detector: {score['detector']:.2%}
|
| 112 |
重复度: {score['repetition']:.2%}
|
| 113 |
-
句式稳定性: {1-score['variance']:.2%}
|
| 114 |
|
| 115 |
{p[:400]}
|
| 116 |
"""
|
| 117 |
)
|
| 118 |
|
| 119 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 120 |
|
| 121 |
return f"""
|
| 122 |
# 综合AI风险率: {overall:.2%}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 123 |
|
| 124 |
-
|
|
|
|
|
|
|
| 125 |
|
| 126 |
---
|
| 127 |
|
| 128 |
""" + "\n\n---\n\n".join(results)
|
| 129 |
|
|
|
|
| 130 |
demo = gr.Interface(
|
| 131 |
fn=analyze_pdf,
|
| 132 |
inputs=gr.File(file_types=[".pdf"]),
|
| 133 |
outputs="markdown",
|
| 134 |
title="论文AIGC风险检测系统",
|
| 135 |
-
description="
|
| 136 |
)
|
| 137 |
|
| 138 |
-
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import os
|
|
|
|
|
|
|
| 3 |
import re
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
from typing import Dict, List
|
| 6 |
|
| 7 |
+
import fitz
|
| 8 |
+
import gradio as gr
|
| 9 |
+
import numpy as np
|
| 10 |
+
import torch
|
| 11 |
+
from transformers import AutoModelForSequenceClassification, AutoTokenizer
|
| 12 |
|
| 13 |
MODEL_NAME = "yuchuantian/AIGC_detector_zhv2"
|
| 14 |
|
| 15 |
+
MAX_PAGES = 120
|
| 16 |
+
MAX_PARAGRAPHS = 120
|
| 17 |
+
MIN_PARAGRAPH_CHARS = 80
|
| 18 |
+
WINDOW_MAX_LENGTH = 512
|
| 19 |
+
WINDOW_STRIDE = 128
|
| 20 |
+
|
| 21 |
+
CALIBRATION_PATH = Path("calibration/model.json")
|
| 22 |
+
|
| 23 |
+
|
| 24 |
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
|
| 25 |
model = AutoModelForSequenceClassification.from_pretrained(MODEL_NAME)
|
| 26 |
+
model.eval()
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def load_calibration_model() -> Dict:
|
| 30 |
+
if not CALIBRATION_PATH.exists():
|
| 31 |
+
return {}
|
| 32 |
+
|
| 33 |
+
try:
|
| 34 |
+
data = json.loads(CALIBRATION_PATH.read_text(encoding="utf-8"))
|
| 35 |
+
except Exception:
|
| 36 |
+
return {}
|
| 37 |
+
|
| 38 |
+
if data.get("model_type") != "linear":
|
| 39 |
+
return {}
|
| 40 |
|
| 41 |
+
required = {"feature_order", "coef", "intercept"}
|
| 42 |
+
if not required.issubset(data.keys()):
|
| 43 |
+
return {}
|
| 44 |
+
|
| 45 |
+
return data
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
CALIBRATION_MODEL = load_calibration_model()
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def normalize_text(text: str) -> str:
|
| 52 |
+
text = re.sub(r"\s+", " ", text)
|
| 53 |
+
return text.strip()
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def split_paragraphs(text: str) -> List[str]:
|
| 57 |
+
paragraphs = re.split(r"\n{2,}", text)
|
| 58 |
+
return [p.strip() for p in paragraphs if len(normalize_text(p)) >= MIN_PARAGRAPH_CHARS]
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def should_skip_paragraph(text: str) -> bool:
|
| 62 |
+
t = normalize_text(text)
|
| 63 |
+
if not t:
|
| 64 |
+
return True
|
| 65 |
+
|
| 66 |
+
if re.search(r"(参考文献|致谢|附录|作者简介)", t[:40], flags=re.IGNORECASE):
|
| 67 |
+
return True
|
| 68 |
+
|
| 69 |
+
cn_chars = len(re.findall(r"[\u4e00-\u9fff]", t))
|
| 70 |
+
digit_punc = len(re.findall(r"[\d\W_]", t))
|
| 71 |
+
if cn_chars < 20 or digit_punc > len(t) * 0.75:
|
| 72 |
+
return True
|
| 73 |
+
|
| 74 |
+
return False
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def extract_pdf_text(pdf_file) -> str:
|
| 78 |
+
doc = fitz.open(pdf_file.name)
|
| 79 |
+
chunks = []
|
| 80 |
|
| 81 |
+
for page_idx, page in enumerate(doc):
|
| 82 |
+
if page_idx >= MAX_PAGES:
|
| 83 |
+
break
|
|
|
|
|
|
|
| 84 |
|
| 85 |
+
blocks = page.get_text("blocks")
|
| 86 |
+
blocks_sorted = sorted(blocks, key=lambda b: (round(b[1], 1), round(b[0], 1)))
|
| 87 |
+
page_text = "\n".join(b[4].strip() for b in blocks_sorted if b[4].strip())
|
| 88 |
|
| 89 |
+
if page_text:
|
| 90 |
+
chunks.append(page_text)
|
| 91 |
|
| 92 |
+
return "\n\n".join(chunks)
|
| 93 |
|
|
|
|
| 94 |
|
| 95 |
+
def calc_repetition(text: str) -> float:
|
| 96 |
+
t = normalize_text(text)
|
| 97 |
+
if not t:
|
| 98 |
+
return 0.0
|
| 99 |
|
| 100 |
+
grams = [t[i : i + 2] for i in range(len(t) - 1)]
|
| 101 |
+
if not grams:
|
| 102 |
+
return 0.0
|
| 103 |
+
|
| 104 |
+
unique_ratio = len(set(grams)) / len(grams)
|
| 105 |
+
return max(0.0, 1.0 - unique_ratio)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def calc_sentence_variance(text: str) -> float:
|
| 109 |
+
sents = re.split(r"[。!?!?]", text)
|
| 110 |
lens = [len(s.strip()) for s in sents if len(s.strip()) > 0]
|
| 111 |
|
| 112 |
if len(lens) < 2:
|
| 113 |
+
return 0.0
|
| 114 |
+
|
| 115 |
+
return float(min(np.var(lens) / 900.0, 1.0))
|
| 116 |
|
|
|
|
| 117 |
|
| 118 |
+
def detector_score(text: str) -> float:
|
| 119 |
inputs = tokenizer(
|
| 120 |
text,
|
|
|
|
| 121 |
truncation=True,
|
| 122 |
+
max_length=WINDOW_MAX_LENGTH,
|
| 123 |
+
stride=WINDOW_STRIDE,
|
| 124 |
+
return_overflowing_tokens=True,
|
| 125 |
+
padding=True,
|
| 126 |
+
return_tensors="pt",
|
| 127 |
)
|
| 128 |
|
| 129 |
with torch.no_grad():
|
| 130 |
outputs = model(**inputs)
|
| 131 |
|
| 132 |
probs = torch.softmax(outputs.logits, dim=-1)
|
| 133 |
+
ai_probs = probs[:, 1].cpu().numpy()
|
| 134 |
|
| 135 |
+
return float(0.75 * np.mean(ai_probs) + 0.25 * np.max(ai_probs))
|
| 136 |
|
|
|
|
|
|
|
| 137 |
|
| 138 |
+
def analyze_paragraph(text: str):
|
| 139 |
+
detector = detector_score(text)
|
| 140 |
repetition = calc_repetition(text)
|
|
|
|
| 141 |
variance = calc_sentence_variance(text)
|
| 142 |
|
| 143 |
+
stable_style = 1 - variance
|
| 144 |
+
risk = detector * 0.78 + repetition * 0.12 + stable_style * 0.10
|
| 145 |
+
risk = float(min(max(risk, 0.0), 1.0))
|
|
|
|
|
|
|
| 146 |
|
| 147 |
return {
|
| 148 |
"detector": detector,
|
| 149 |
"repetition": repetition,
|
| 150 |
"variance": variance,
|
| 151 |
+
"risk": risk,
|
| 152 |
}
|
| 153 |
|
|
|
|
|
|
|
| 154 |
|
| 155 |
+
def clip01(v: float) -> float:
|
| 156 |
+
return float(min(max(v, 0.0), 1.0))
|
| 157 |
+
|
| 158 |
+
|
| 159 |
+
def build_doc_features(risks: List[float]) -> Dict[str, float]:
|
| 160 |
+
arr = np.array(risks, dtype=float)
|
| 161 |
+
return {
|
| 162 |
+
"overall": clip01(float(np.mean(arr))),
|
| 163 |
+
"p90": clip01(float(np.percentile(arr, 90))),
|
| 164 |
+
"high_ratio": clip01(float(np.mean(arr > 0.75))),
|
| 165 |
+
"mid_ratio": clip01(float(np.mean((arr > 0.55) & (arr <= 0.75)))),
|
| 166 |
+
"std": clip01(float(np.std(arr))),
|
| 167 |
+
}
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
def predict_kn_like_rate(features: Dict[str, float]) -> float:
|
| 171 |
+
if not CALIBRATION_MODEL:
|
| 172 |
+
return features["overall"]
|
| 173 |
+
|
| 174 |
+
order = CALIBRATION_MODEL["feature_order"]
|
| 175 |
+
coef = CALIBRATION_MODEL["coef"]
|
| 176 |
+
intercept = float(CALIBRATION_MODEL["intercept"])
|
| 177 |
|
| 178 |
+
x = np.array([features.get(name, 0.0) for name in order], dtype=float)
|
| 179 |
+
y = float(np.dot(x, np.array(coef, dtype=float)) + intercept)
|
| 180 |
|
| 181 |
+
return clip01(y)
|
| 182 |
+
|
| 183 |
+
|
| 184 |
+
def analyze_pdf(pdf_file):
|
| 185 |
+
text = extract_pdf_text(pdf_file)
|
| 186 |
paragraphs = split_paragraphs(text)
|
| 187 |
+
paragraphs = [p for p in paragraphs if not should_skip_paragraph(p)]
|
| 188 |
+
paragraphs = paragraphs[:MAX_PARAGRAPHS]
|
| 189 |
|
| 190 |
+
if not paragraphs:
|
| 191 |
+
return "未提取到可分析正文。请尝试文本层可复制的 PDF,或调整排版后再试。"
|
| 192 |
|
| 193 |
+
results = []
|
| 194 |
risks = []
|
| 195 |
|
| 196 |
+
for p in paragraphs:
|
| 197 |
score = analyze_paragraph(p)
|
|
|
|
| 198 |
risks.append(score["risk"])
|
| 199 |
|
| 200 |
level = "🟢"
|
|
|
|
| 201 |
if score["risk"] > 0.75:
|
| 202 |
level = "🔴"
|
| 203 |
elif score["risk"] > 0.55:
|
|
|
|
| 209 |
|
| 210 |
Detector: {score['detector']:.2%}
|
| 211 |
重复度: {score['repetition']:.2%}
|
| 212 |
+
句式稳定性: {1 - score['variance']:.2%}
|
| 213 |
|
| 214 |
{p[:400]}
|
| 215 |
"""
|
| 216 |
)
|
| 217 |
|
| 218 |
+
features = build_doc_features(risks)
|
| 219 |
+
overall = features["overall"]
|
| 220 |
+
high_ratio = features["high_ratio"]
|
| 221 |
+
mid_ratio = features["mid_ratio"]
|
| 222 |
+
kn_like_rate = predict_kn_like_rate(features)
|
| 223 |
+
|
| 224 |
+
mode_line = "当前模式: 原始风险率(未加载校准模型)"
|
| 225 |
+
if CALIBRATION_MODEL:
|
| 226 |
+
mode_line = "当前模式: 知网对齐预测率(已加载校准模型)"
|
| 227 |
|
| 228 |
return f"""
|
| 229 |
# 综合AI风险率: {overall:.2%}
|
| 230 |
+
# 预测知网AIGC率: {kn_like_rate:.2%}
|
| 231 |
+
高风险段落占比: {high_ratio:.2%}
|
| 232 |
+
中风险段落占比: {mid_ratio:.2%}
|
| 233 |
+
有效段落数: {len(paragraphs)}
|
| 234 |
|
| 235 |
+
{mode_line}
|
| 236 |
+
|
| 237 |
+
(说明:该结果为“风险分析与校准预测”,并非官方系统结果)
|
| 238 |
|
| 239 |
---
|
| 240 |
|
| 241 |
""" + "\n\n---\n\n".join(results)
|
| 242 |
|
| 243 |
+
|
| 244 |
demo = gr.Interface(
|
| 245 |
fn=analyze_pdf,
|
| 246 |
inputs=gr.File(file_types=[".pdf"]),
|
| 247 |
outputs="markdown",
|
| 248 |
title="论文AIGC风险检测系统",
|
| 249 |
+
description="论文AIGC风险分析(含知网对齐预测)",
|
| 250 |
)
|
| 251 |
|
| 252 |
+
|
| 253 |
+
demo.launch()
|
calibration/model.json
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"model_type": "linear",
|
| 3 |
+
"feature_order": ["overall", "p90", "high_ratio", "mid_ratio", "std"],
|
| 4 |
+
"coef": [1.0, 0.0, 0.0, 0.0, 0.0],
|
| 5 |
+
"intercept": 0.0,
|
| 6 |
+
"note": "default passthrough calibration; replace with trained calibration model"
|
| 7 |
+
}
|
scripts/train_calibration.py
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import argparse
|
| 2 |
+
import json
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
import pandas as pd
|
| 7 |
+
from sklearn.linear_model import LinearRegression
|
| 8 |
+
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score
|
| 9 |
+
|
| 10 |
+
FEATURE_ORDER = ["overall", "p90", "high_ratio", "mid_ratio", "std"]
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def clip01(x):
|
| 14 |
+
return np.clip(x, 0.0, 1.0)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def main():
|
| 18 |
+
parser = argparse.ArgumentParser(description="Train kn-like calibration model")
|
| 19 |
+
parser.add_argument("--input", required=True, help="CSV with feature columns + target")
|
| 20 |
+
parser.add_argument("--target", default="kn_rate", help="target column name")
|
| 21 |
+
parser.add_argument("--out", default="calibration/model.json", help="output model json")
|
| 22 |
+
args = parser.parse_args()
|
| 23 |
+
|
| 24 |
+
df = pd.read_csv(args.input)
|
| 25 |
+
|
| 26 |
+
missing = [c for c in FEATURE_ORDER + [args.target] if c not in df.columns]
|
| 27 |
+
if missing:
|
| 28 |
+
raise ValueError(f"Missing columns: {missing}")
|
| 29 |
+
|
| 30 |
+
X = df[FEATURE_ORDER].astype(float).values
|
| 31 |
+
y = df[args.target].astype(float).values
|
| 32 |
+
|
| 33 |
+
model = LinearRegression()
|
| 34 |
+
model.fit(X, y)
|
| 35 |
+
|
| 36 |
+
pred = clip01(model.predict(X))
|
| 37 |
+
|
| 38 |
+
metrics = {
|
| 39 |
+
"mae": float(mean_absolute_error(y, pred)),
|
| 40 |
+
"rmse": float(np.sqrt(mean_squared_error(y, pred))),
|
| 41 |
+
"r2": float(r2_score(y, pred)),
|
| 42 |
+
"n": int(len(df)),
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
payload = {
|
| 46 |
+
"model_type": "linear",
|
| 47 |
+
"feature_order": FEATURE_ORDER,
|
| 48 |
+
"coef": [float(v) for v in model.coef_.tolist()],
|
| 49 |
+
"intercept": float(model.intercept_),
|
| 50 |
+
"train_metrics": metrics,
|
| 51 |
+
}
|
| 52 |
+
|
| 53 |
+
out = Path(args.out)
|
| 54 |
+
out.parent.mkdir(parents=True, exist_ok=True)
|
| 55 |
+
out.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
|
| 56 |
+
|
| 57 |
+
print("Saved:", out)
|
| 58 |
+
print("Metrics:", metrics)
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
if __name__ == "__main__":
|
| 62 |
+
main()
|