-
Notifications
You must be signed in to change notification settings - Fork 7
Expand file tree
/
Copy pathquickstart.py
More file actions
56 lines (44 loc) · 1.68 KB
/
Copy pathquickstart.py
File metadata and controls
56 lines (44 loc) · 1.68 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
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import torch
from sklearn.datasets import load_breast_cancer
from torch import Tensor
import sdm
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
df = load_breast_cancer(as_frame=True).frame
table = sdm.TableTensor.from_pandas(
df=df,
stypes=sdm.infer_stypes(df, overrides={"target": "categorical"}),
device=device,
)
model = sdm.models.TabICLv2(device=device)
# Default in-context learning forward pass:
with torch.amp.autocast(device.type, torch.float16, enabled=table.is_cuda):
model(
x_context=table[:300].drop_columns("target"),
y_context=table[:300, "target"],
x_query=table[300:].drop_columns("target"),
num_estimators=2,
)
# Fit + Predict forward pass via key/value caching for fast inference:
with torch.amp.autocast(device.type, torch.float16, enabled=table.is_cuda):
model.fit(
x=table[:300].drop_columns("target"),
y=table[:300, "target"],
num_estimators=2,
)
model.predict(table[300:].drop_columns("target"))
model.clear()
# Capturing embeddings:
embeddings: list[Tensor] = []
def _embedding(module: torch.nn.Module, args: tuple[Tensor, ...]) -> None:
embeddings.append(args[0])
head = model.models["classification"].icl_block.head
handle = head.register_forward_pre_hook(_embedding)
with torch.amp.autocast(device.type, torch.float16, enabled=table.is_cuda):
model(
x_context=table[:300].drop_columns("target"),
y_context=table[:300, "target"],
x_query=table[300:].drop_columns("target"),
)
handle.remove()