Repository navigation
Expand file tree
/
Copy pathmain.py
More file actions
171 lines (142 loc) · 5.17 KB
/
Copy pathmain.py
File metadata and controls
171 lines (142 loc) · 5.17 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
import argparse
import sys
import joblib
from src.dataset import load_dataset, preprocess_dataset
from src.train_model import (
train_clickbait_model,
print_metrics,
save_model,
compute_classification_metrics,
)
from src.predict import predict_headline
from src.features import transform_to_dataframe
DEFAULT_DATA_PATH = "data/raw/clickbait_data.csv"
DEFAULT_MODEL_PATH = "models/clickbait_lr_tfidf_v1.0.pkl"
DEFAULT_VECTORIZER_PATH = "models/clickbait_tfidf_vectorizer_v1.0.pkl"
def train_command(
data_path: str = DEFAULT_DATA_PATH,
model_path: str = DEFAULT_MODEL_PATH,
vectorizer_path: str = DEFAULT_VECTORIZER_PATH,
) -> None:
model, vectorizer, metrics = train_clickbait_model(data_path)
print_metrics(metrics)
save_model(model, vectorizer, model_path=model_path, vectorizer_path=vectorizer_path)
print(f"Model saved to: {model_path}")
print(f"Vectorizer saved to: {vectorizer_path}")
def predict_command(
headline: str,
model_path: str = DEFAULT_MODEL_PATH,
vectorizer_path: str = DEFAULT_VECTORIZER_PATH,
) -> None:
try:
model = joblib.load(model_path)
vectorizer = joblib.load(vectorizer_path)
except FileNotFoundError as exc:
print(f"Error: required file not found: {exc}", file=sys.stderr)
sys.exit(1)
result = predict_headline(
headline, model, vectorizer, return_probability=True
)
label = result["label"]
probability_clickbait = result["probability_clickbait"]
print(f"Headline: {headline}")
print(f"Prediction: {label.upper()}")
print(f"Probability (clickbait): {probability_clickbait:.4f}")
def metrics_command(
data_path: str = DEFAULT_DATA_PATH,
model_path: str = DEFAULT_MODEL_PATH,
vectorizer_path: str = DEFAULT_VECTORIZER_PATH,
) -> None:
try:
model = joblib.load(model_path)
vectorizer = joblib.load(vectorizer_path)
except FileNotFoundError as exc:
print(f"Error: required file not found: {exc}", file=sys.stderr)
sys.exit(1)
df = load_dataset(data_path)
df = preprocess_dataset(df)
X = df["clean_headline"]
y = df["clickbait"]
X_tfidf = transform_to_dataframe(vectorizer, X, index=df.index)
y_pred = model.predict(X_tfidf)
metrics = compute_classification_metrics(y, y_pred)
print_metrics(metrics)
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="ClickShield: train and run a clickbait detection model."
)
subparsers = parser.add_subparsers(dest="command", required=True)
train_parser = subparsers.add_parser("train", help="Train a new clickbait model.")
train_parser.add_argument(
"--data-path",
default=DEFAULT_DATA_PATH,
help=f"Path to training data CSV (default: {DEFAULT_DATA_PATH})",
)
train_parser.add_argument(
"--model-path",
default=DEFAULT_MODEL_PATH,
help=f"Path to save trained model (default: {DEFAULT_MODEL_PATH})",
)
train_parser.add_argument(
"--vectorizer-path",
default=DEFAULT_VECTORIZER_PATH,
help=f"Path to save vectorizer (default: {DEFAULT_VECTORIZER_PATH})",
)
predict_parser = subparsers.add_parser(
"predict", help="Predict whether a headline is clickbait."
)
predict_parser.add_argument("headline", help="Headline text to classify.")
predict_parser.add_argument(
"--model-path",
default=DEFAULT_MODEL_PATH,
help=f"Path to trained model (default: {DEFAULT_MODEL_PATH})",
)
predict_parser.add_argument(
"--vectorizer-path",
default=DEFAULT_VECTORIZER_PATH,
help=f"Path to trained vectorizer (default: {DEFAULT_VECTORIZER_PATH})",
)
metrics_parser = subparsers.add_parser(
"metrics",
help="Compute and print evaluation metrics using an existing model and dataset.",
)
metrics_parser.add_argument(
"--data-path",
default=DEFAULT_DATA_PATH,
help=f"Path to evaluation data CSV (default: {DEFAULT_DATA_PATH})",
)
metrics_parser.add_argument(
"--model-path",
default=DEFAULT_MODEL_PATH,
help=f"Path to trained model (default: {DEFAULT_MODEL_PATH})",
)
metrics_parser.add_argument(
"--vectorizer-path",
default=DEFAULT_VECTORIZER_PATH,
help=f"Path to trained vectorizer (default: {DEFAULT_VECTORIZER_PATH})",
)
return parser.parse_args(argv)
def main(argv: list[str] | None = None) -> None:
args = parse_args(argv)
if args.command == "train":
train_command(
data_path=args.data_path,
model_path=args.model_path,
vectorizer_path=args.vectorizer_path,
)
elif args.command == "predict":
predict_command(
headline=args.headline,
model_path=args.model_path,
vectorizer_path=args.vectorizer_path,
)
elif args.command == "metrics":
metrics_command(
data_path=args.data_path,
model_path=args.model_path,
vectorizer_path=args.vectorizer_path,
)
else:
raise AssertionError("Unhandled command")
if __name__ == "__main__":
main()