-
Notifications
You must be signed in to change notification settings - Fork 16
Expand file tree
/
Copy pathtext.py
More file actions
76 lines (63 loc) · 2.24 KB
/
Copy pathtext.py
File metadata and controls
76 lines (63 loc) · 2.24 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
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""TabICLv2 on STRABLE example.
Predict CLEAR Corpus readability with TabICLv2 and optional text features.
Text features are either embeded via a Sentence Transformer followed by PCA, or
encoded via character n-gram TF-IDF.
"""
import argparse
import pyarrow.parquet as pq
import torch
from huggingface_hub import hf_hub_download
import sdm
import sdm.processing as sp
parser = argparse.ArgumentParser()
parser.add_argument(
"--text-processor",
choices=("embed", "tfidf", "none"),
default="embed",
)
parser.add_argument("--seed", type=int, default=0)
args = parser.parse_args()
torch.manual_seed(args.seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
data_path = hf_hub_download(
repo_id="inria-soda/STRABLE-benchmark",
filename="clear-corpus/data.parquet",
repo_type="dataset",
)
target_name = "BT Easiness"
arrow_table = pq.read_table(data_path)
table = sdm.TableTensor.from_arrow(
table=arrow_table,
stypes=sdm.infer_stypes(
arrow_table,
text="off" if args.text_processor == "none" else "infer",
),
device=device,
)
table = table[torch.randperm(table.size(0), device=device)]
context, query = table.split(int(0.8 * table.size(0)))
model = sdm.models.TabICLv2(device=device)
recipe = model.default_recipe()
if args.text_processor != "none":
if args.text_processor == "tfidf":
text_processor = sp.TFIDF(ngram_range=(4, 6), max_features=256)
else:
text_processor = [
sp.SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2"),
sp.PCA(64),
]
recipe.prepend_features(sp.StypeDispatch(text=text_processor))
with torch.amp.autocast(device.type, torch.float16, enabled=table.is_cuda):
pred = model(
x_context=context.drop_columns(target_name),
y_context=context[:, target_name],
x_query=query.drop_columns(target_name),
recipe=recipe,
num_estimators=8,
).numerical.mean(dim=-1)
y_query = query[:, target_name].numerical.squeeze(-1)
rmse = (pred - y_query).pow(2).mean().sqrt()
mae = (pred - y_query).abs().mean()
print(f"RMSE: {rmse:.3f}, MAE: {mae:.3f}")