tensorless 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,179 @@
1
+ """Tabular preprocessing.
2
+
3
+ Fits simple, fully-reversible preprocessing on tabular records:
4
+ - numeric columns -> standardized (mean/std), missing values imputed
5
+ with the training-set mean
6
+ - categorical columns -> integer-indexed vocabulary (+ <unk>/<missing>),
7
+ fed into per-column embeddings by the MLP model
8
+
9
+ The fitted state is small and JSON-serializable, so it can be embedded
10
+ directly inside a `.tl` file and reproduced exactly at inference time.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ from dataclasses import dataclass, field
16
+ from typing import Any, Dict, List, Optional, Tuple
17
+
18
+ import torch
19
+
20
+ _MISSING = "<missing>"
21
+ _UNK = "<unk>"
22
+
23
+
24
+ def _try_float(v: Any) -> Optional[float]:
25
+ if v is None:
26
+ return None
27
+ if isinstance(v, str) and v.strip() == "":
28
+ return None
29
+ try:
30
+ return float(v)
31
+ except (TypeError, ValueError):
32
+ return None
33
+
34
+
35
+ @dataclass
36
+ class ColumnStats:
37
+ kind: str # "numeric" | "categorical"
38
+ mean: float = 0.0
39
+ std: float = 1.0
40
+ vocab: List[str] = field(default_factory=list)
41
+
42
+
43
+ @dataclass
44
+ class TabularPreprocessor:
45
+ feature_columns: List[str] = field(default_factory=list)
46
+ target_column: Optional[str] = None
47
+ task: str = "classification"
48
+ column_stats: Dict[str, ColumnStats] = field(default_factory=dict)
49
+ classes: List[str] = field(default_factory=list) # for classification targets
50
+ target_mean: float = 0.0
51
+ target_std: float = 1.0
52
+
53
+ @property
54
+ def numeric_columns(self) -> List[str]:
55
+ return [c for c in self.feature_columns if self.column_stats[c].kind == "numeric"]
56
+
57
+ @property
58
+ def categorical_columns(self) -> List[str]:
59
+ return [c for c in self.feature_columns if self.column_stats[c].kind == "categorical"]
60
+
61
+ def fit(
62
+ self,
63
+ records: List[Dict[str, Any]],
64
+ columns: List[str],
65
+ target_column: str,
66
+ task: str,
67
+ ) -> "TabularPreprocessor":
68
+ self.target_column = target_column
69
+ self.task = task
70
+ self.feature_columns = [c for c in columns if c != target_column]
71
+
72
+ for col in self.feature_columns:
73
+ values = [r.get(col) for r in records]
74
+ numeric_vals = [_try_float(v) for v in values]
75
+ n_present = sum(1 for v in values if v not in (None, ""))
76
+ n_numeric = sum(1 for v in numeric_vals if v is not None)
77
+ if n_present > 0 and n_numeric / n_present > 0.95:
78
+ nums = [v for v in numeric_vals if v is not None]
79
+ mean = sum(nums) / len(nums) if nums else 0.0
80
+ var = sum((x - mean) ** 2 for x in nums) / len(nums) if nums else 1.0
81
+ std = max(var ** 0.5, 1e-6)
82
+ self.column_stats[col] = ColumnStats(kind="numeric", mean=mean, std=std)
83
+ else:
84
+ cats = sorted({str(v) for v in values if v not in (None, "")})
85
+ vocab = [_MISSING, _UNK] + cats
86
+ self.column_stats[col] = ColumnStats(kind="categorical", vocab=vocab)
87
+
88
+ target_vals = [r.get(target_column) for r in records]
89
+ if task == "regression":
90
+ nums = [v for v in (_try_float(v) for v in target_vals) if v is not None]
91
+ self.target_mean = sum(nums) / len(nums) if nums else 0.0
92
+ var = sum((x - self.target_mean) ** 2 for x in nums) / len(nums) if nums else 1.0
93
+ self.target_std = max(var ** 0.5, 1e-6)
94
+ else:
95
+ self.classes = sorted({str(v) for v in target_vals if v not in (None, "")})
96
+
97
+ return self
98
+
99
+ def transform(
100
+ self, records: List[Dict[str, Any]], with_target: bool = True
101
+ ) -> Dict[str, torch.Tensor]:
102
+ n = len(records)
103
+ num_cols = self.numeric_columns
104
+ cat_cols = self.categorical_columns
105
+
106
+ numeric = torch.zeros(n, max(len(num_cols), 1), dtype=torch.float32)
107
+ categorical = torch.zeros(n, max(len(cat_cols), 1), dtype=torch.long)
108
+
109
+ for i, r in enumerate(records):
110
+ for j, col in enumerate(num_cols):
111
+ stats = self.column_stats[col]
112
+ v = _try_float(r.get(col))
113
+ if v is None:
114
+ v = stats.mean
115
+ numeric[i, j] = (v - stats.mean) / stats.std
116
+ for j, col in enumerate(cat_cols):
117
+ stats = self.column_stats[col]
118
+ raw = r.get(col)
119
+ key = _MISSING if raw in (None, "") else str(raw)
120
+ idx = stats.vocab.index(key) if key in stats.vocab else stats.vocab.index(_UNK)
121
+ categorical[i, j] = idx
122
+
123
+ out = {"numeric": numeric, "categorical": categorical}
124
+
125
+ if with_target and self.target_column is not None:
126
+ if self.task == "regression":
127
+ target = torch.zeros(n, dtype=torch.float32)
128
+ for i, r in enumerate(records):
129
+ v = _try_float(r.get(self.target_column))
130
+ v = self.target_mean if v is None else v
131
+ target[i] = (v - self.target_mean) / self.target_std
132
+ out["target"] = target
133
+ else:
134
+ target = torch.zeros(n, dtype=torch.long)
135
+ for i, r in enumerate(records):
136
+ raw = str(r.get(self.target_column))
137
+ idx = self.classes.index(raw) if raw in self.classes else 0
138
+ target[i] = idx
139
+ out["target"] = target
140
+
141
+ return out
142
+
143
+ def inverse_target(self, values: torch.Tensor) -> List[Any]:
144
+ if self.task == "regression":
145
+ return [(v.item() * self.target_std + self.target_mean) for v in values]
146
+ return [self.classes[int(v.item())] for v in values]
147
+
148
+ def categorical_vocab_sizes(self) -> List[int]:
149
+ return [len(self.column_stats[c].vocab) for c in self.categorical_columns]
150
+
151
+ def state_dict(self) -> Dict:
152
+ return {
153
+ "feature_columns": self.feature_columns,
154
+ "target_column": self.target_column,
155
+ "task": self.task,
156
+ "column_stats": {
157
+ c: {"kind": s.kind, "mean": s.mean, "std": s.std, "vocab": s.vocab}
158
+ for c, s in self.column_stats.items()
159
+ },
160
+ "classes": self.classes,
161
+ "target_mean": self.target_mean,
162
+ "target_std": self.target_std,
163
+ }
164
+
165
+ @classmethod
166
+ def from_state_dict(cls, state: Dict) -> "TabularPreprocessor":
167
+ prep = cls(
168
+ feature_columns=list(state["feature_columns"]),
169
+ target_column=state["target_column"],
170
+ task=state["task"],
171
+ classes=list(state.get("classes", [])),
172
+ target_mean=state.get("target_mean", 0.0),
173
+ target_std=state.get("target_std", 1.0),
174
+ )
175
+ prep.column_stats = {
176
+ c: ColumnStats(kind=v["kind"], mean=v.get("mean", 0.0), std=v.get("std", 1.0), vocab=list(v.get("vocab", [])))
177
+ for c, v in state["column_stats"].items()
178
+ }
179
+ return prep
@@ -0,0 +1,3 @@
1
+ from .device import auto_select_device, get_torch_device
2
+
3
+ __all__ = ["auto_select_device", "get_torch_device"]
@@ -0,0 +1,107 @@
1
+ """Hardware auto-detection.
2
+
3
+ Preference order is TPU -> GPU -> CPU, but selection is "intelligent"
4
+ rather than blind: we verify each backend is actually usable (not just
5
+ importable) before choosing it, and we fall back gracefully -- including
6
+ at *runtime*, if a chosen device turns out to error out mid-training, the
7
+ trainer (see `training/trainer.py`) will catch that and fall back too.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from typing import Optional, Tuple
13
+
14
+ import torch
15
+
16
+
17
+ def _tpu_available() -> bool:
18
+ try:
19
+ import torch_xla.core.xla_model as xm # noqa: F401
20
+
21
+ return True
22
+ except Exception:
23
+ return False
24
+
25
+
26
+ def _cuda_available() -> bool:
27
+ try:
28
+ return torch.cuda.is_available() and torch.cuda.device_count() > 0
29
+ except Exception:
30
+ return False
31
+
32
+
33
+ def _mps_available() -> bool:
34
+ try:
35
+ return torch.backends.mps.is_available()
36
+ except Exception:
37
+ return False
38
+
39
+
40
+ def _cuda_supports_bf16() -> bool:
41
+ try:
42
+ return torch.cuda.is_bf16_supported()
43
+ except Exception:
44
+ return False
45
+
46
+
47
+ def auto_select_device(user_device: Optional[str], user_precision: Optional[str]) -> Tuple[str, str]:
48
+ """Resolve the device and precision to use.
49
+
50
+ `user_device` / `user_precision` are honored if given (with a
51
+ graceful downgrade if the requested device isn't actually available).
52
+ Otherwise we pick automatically: tpu > cuda > mps > cpu.
53
+ """
54
+ if user_device is not None:
55
+ device = user_device
56
+ if device == "tpu" and not _tpu_available():
57
+ device = "cuda" if _cuda_available() else "cpu"
58
+ elif device == "cuda" and not _cuda_available():
59
+ device = "cpu"
60
+ elif device == "mps" and not _mps_available():
61
+ device = "cpu"
62
+ else:
63
+ if _tpu_available():
64
+ device = "tpu"
65
+ elif _cuda_available():
66
+ device = "cuda"
67
+ elif _mps_available():
68
+ device = "mps"
69
+ else:
70
+ device = "cpu"
71
+
72
+ if user_precision is not None:
73
+ precision = user_precision
74
+ else:
75
+ if device == "cuda" and _cuda_supports_bf16():
76
+ precision = "bf16"
77
+ elif device == "cuda":
78
+ precision = "fp16"
79
+ elif device == "tpu":
80
+ precision = "bf16"
81
+ else:
82
+ # CPU and MPS: stick to fp32 for correctness/stability by default.
83
+ precision = "fp32"
84
+
85
+ return device, precision
86
+
87
+
88
+ def get_torch_device(device: str) -> torch.device:
89
+ """Convert our string device name into a torch.device, with a
90
+ runtime fallback to CPU if the requested backend is unavailable.
91
+ """
92
+ try:
93
+ if device == "tpu":
94
+ import torch_xla.core.xla_model as xm
95
+
96
+ return xm.xla_device()
97
+ if device == "cuda":
98
+ if not _cuda_available():
99
+ return torch.device("cpu")
100
+ return torch.device("cuda")
101
+ if device == "mps":
102
+ if not _mps_available():
103
+ return torch.device("cpu")
104
+ return torch.device("mps")
105
+ return torch.device("cpu")
106
+ except Exception:
107
+ return torch.device("cpu")
tensorless/errors.py ADDED
@@ -0,0 +1,36 @@
1
+ """Tensorless error hierarchy.
2
+
3
+ All errors raised by Tensorless inherit from :class:`TensorlessError`, so
4
+ callers can do::
5
+
6
+ try:
7
+ tl.train("./data")
8
+ except tl.TensorlessError as e:
9
+ ...
10
+
11
+ instead of catching bare Exception.
12
+ """
13
+
14
+
15
+ class TensorlessError(Exception):
16
+ """Base class for all Tensorless errors."""
17
+
18
+
19
+ class DataError(TensorlessError):
20
+ """Raised when a dataset cannot be read, is empty, or is malformed."""
21
+
22
+
23
+ class ConfigError(TensorlessError):
24
+ """Raised when user-supplied configuration is invalid or contradictory."""
25
+
26
+
27
+ class ModelError(TensorlessError):
28
+ """Raised for unsupported model/task combinations or model build failures."""
29
+
30
+
31
+ class CheckpointError(TensorlessError):
32
+ """Raised when a checkpoint is missing, corrupt, or incompatible."""
33
+
34
+
35
+ class SerializationError(TensorlessError):
36
+ """Raised when a `.tl` file cannot be written or read."""
@@ -0,0 +1,5 @@
1
+ from .registry import build_model
2
+ from .transformer import TinyTransformer
3
+ from .mlp import TabularMLP
4
+
5
+ __all__ = ["build_model", "TinyTransformer", "TabularMLP"]
@@ -0,0 +1,64 @@
1
+ """MLP model for tabular data (classification / regression).
2
+
3
+ Numeric columns feed directly into the network; each categorical column
4
+ gets its own small embedding table, and all features are concatenated
5
+ before the hidden layers.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from typing import List
11
+
12
+ import torch
13
+ import torch.nn as nn
14
+ import torch.nn.functional as F
15
+
16
+
17
+ def _embedding_dim(vocab_size: int) -> int:
18
+ # Common rule of thumb, capped to keep tiny-data models tiny.
19
+ return max(2, min(32, round(1.6 * (vocab_size ** 0.56))))
20
+
21
+
22
+ class TabularMLP(nn.Module):
23
+ def __init__(
24
+ self,
25
+ n_numeric: int,
26
+ categorical_vocab_sizes: List[int],
27
+ d_model: int,
28
+ layers: int,
29
+ dropout: float,
30
+ task: str,
31
+ n_classes: int = 0,
32
+ ):
33
+ super().__init__()
34
+ self.task = task
35
+ self.n_numeric = n_numeric
36
+ self.categorical_vocab_sizes = categorical_vocab_sizes
37
+
38
+ self.embeddings = nn.ModuleList(
39
+ [nn.Embedding(v, _embedding_dim(v)) for v in categorical_vocab_sizes]
40
+ )
41
+ cat_dim = sum(_embedding_dim(v) for v in categorical_vocab_sizes)
42
+ in_dim = n_numeric + cat_dim
43
+
44
+ dims = [in_dim] + [d_model] * layers
45
+ modules = []
46
+ for i in range(len(dims) - 1):
47
+ modules += [nn.Linear(dims[i], dims[i + 1]), nn.ReLU(), nn.Dropout(dropout)]
48
+ self.backbone = nn.Sequential(*modules)
49
+
50
+ out_dim = n_classes if task == "classification" else 1
51
+ self.head = nn.Linear(dims[-1], out_dim)
52
+
53
+ def forward(self, numeric: torch.Tensor, categorical: torch.Tensor) -> torch.Tensor:
54
+ parts = []
55
+ if self.n_numeric > 0:
56
+ parts.append(numeric)
57
+ for i, emb in enumerate(self.embeddings):
58
+ parts.append(emb(categorical[:, i]))
59
+ x = torch.cat(parts, dim=1) if parts else numeric
60
+ x = self.backbone(x)
61
+ out = self.head(x)
62
+ if self.task == "regression":
63
+ return out.squeeze(-1)
64
+ return out
@@ -0,0 +1,53 @@
1
+ """Model registry.
2
+
3
+ Central place that knows how to build a model given a task + resolved
4
+ config + the dataset-derived metadata (vocab size, n_classes, etc). Kept
5
+ separate from `auto/config.py` so new model types/backends can be added
6
+ here without touching auto-configuration logic, per the "extensible
7
+ architecture" requirement.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from typing import Any, Dict
13
+
14
+ import torch.nn as nn
15
+
16
+ from .transformer import TinyTransformer
17
+ from .mlp import TabularMLP
18
+ from ..errors import ModelError
19
+
20
+
21
+ def build_model(task: str, model_type: str, cfg: Dict[str, Any], meta: Dict[str, Any]) -> nn.Module:
22
+ """Build a fresh, randomly-initialized model.
23
+
24
+ `cfg` is the resolved training config (dict). `meta` carries
25
+ task-specific sizing info produced during data prep, e.g.:
26
+ - text tasks: {"vocab_size": int, "pad_id": int, "n_classes": int}
27
+ - tabular tasks: {"n_numeric": int, "categorical_vocab_sizes": [...], "n_classes": int}
28
+ """
29
+ if model_type == "transformer":
30
+ return TinyTransformer(
31
+ vocab_size=meta["vocab_size"],
32
+ d_model=cfg["d_model"],
33
+ layers=cfg["layers"],
34
+ heads=cfg["heads"],
35
+ ff_mult=cfg["ff_mult"],
36
+ dropout=cfg["dropout"],
37
+ max_seq_len=cfg["max_seq_len"],
38
+ task=task,
39
+ n_classes=meta.get("n_classes", 0),
40
+ pad_id=meta.get("pad_id", 0),
41
+ )
42
+ elif model_type == "mlp":
43
+ return TabularMLP(
44
+ n_numeric=meta["n_numeric"],
45
+ categorical_vocab_sizes=meta["categorical_vocab_sizes"],
46
+ d_model=cfg["d_model"],
47
+ layers=cfg["layers"],
48
+ dropout=cfg["dropout"],
49
+ task=task,
50
+ n_classes=meta.get("n_classes", 0),
51
+ )
52
+ else:
53
+ raise ModelError(f"Unknown model_type '{model_type}'.")
@@ -0,0 +1,175 @@
1
+ """A small, dependency-free (beyond PyTorch) GPT-style decoder transformer.
2
+
3
+ Used for:
4
+ - "text-generation": next-token prediction over the char vocabulary
5
+ - "text-classification": same backbone, with a classification head on
6
+ the final token's hidden state instead of a language-modeling head
7
+
8
+ Kept intentionally compact -- this is not meant to compete with
9
+ production LLM training frameworks, it's meant to give Tensorless a real,
10
+ working, from-scratch model that trains fast enough on CPU for the
11
+ "zero setup" experience to actually be pleasant.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import math
17
+ from typing import Optional
18
+
19
+ import torch
20
+ import torch.nn as nn
21
+ import torch.nn.functional as F
22
+
23
+
24
+ class CausalSelfAttention(nn.Module):
25
+ def __init__(self, d_model: int, heads: int, dropout: float):
26
+ super().__init__()
27
+ assert d_model % heads == 0, "d_model must be divisible by heads"
28
+ self.heads = heads
29
+ self.head_dim = d_model // heads
30
+ self.qkv = nn.Linear(d_model, 3 * d_model)
31
+ self.proj = nn.Linear(d_model, d_model)
32
+ self.dropout = dropout
33
+ self.resid_drop = nn.Dropout(dropout)
34
+
35
+ def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
36
+ B, T, C = x.shape
37
+ qkv = self.qkv(x)
38
+ q, k, v = qkv.split(C, dim=2)
39
+ q = q.view(B, T, self.heads, self.head_dim).transpose(1, 2)
40
+ k = k.view(B, T, self.heads, self.head_dim).transpose(1, 2)
41
+ v = v.view(B, T, self.heads, self.head_dim).transpose(1, 2)
42
+
43
+ out = F.scaled_dot_product_attention(
44
+ q, k, v, attn_mask=attn_mask, dropout_p=self.dropout if self.training else 0.0, is_causal=attn_mask is None
45
+ )
46
+ out = out.transpose(1, 2).contiguous().view(B, T, C)
47
+ return self.resid_drop(self.proj(out))
48
+
49
+
50
+ class MLP(nn.Module):
51
+ def __init__(self, d_model: int, ff_mult: int, dropout: float):
52
+ super().__init__()
53
+ self.fc1 = nn.Linear(d_model, d_model * ff_mult)
54
+ self.fc2 = nn.Linear(d_model * ff_mult, d_model)
55
+ self.drop = nn.Dropout(dropout)
56
+
57
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
58
+ return self.drop(self.fc2(F.gelu(self.fc1(x))))
59
+
60
+
61
+ class Block(nn.Module):
62
+ def __init__(self, d_model: int, heads: int, ff_mult: int, dropout: float):
63
+ super().__init__()
64
+ self.ln1 = nn.LayerNorm(d_model)
65
+ self.attn = CausalSelfAttention(d_model, heads, dropout)
66
+ self.ln2 = nn.LayerNorm(d_model)
67
+ self.mlp = MLP(d_model, ff_mult, dropout)
68
+
69
+ def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
70
+ x = x + self.attn(self.ln1(x), attn_mask=attn_mask)
71
+ x = x + self.mlp(self.ln2(x))
72
+ return x
73
+
74
+
75
+ class TinyTransformer(nn.Module):
76
+ """Decoder-only transformer usable for LM or sequence classification."""
77
+
78
+ def __init__(
79
+ self,
80
+ vocab_size: int,
81
+ d_model: int,
82
+ layers: int,
83
+ heads: int,
84
+ ff_mult: int,
85
+ dropout: float,
86
+ max_seq_len: int,
87
+ task: str = "text-generation",
88
+ n_classes: int = 0,
89
+ pad_id: int = 0,
90
+ ):
91
+ super().__init__()
92
+ self.task = task
93
+ self.max_seq_len = max_seq_len
94
+ self.pad_id = pad_id
95
+
96
+ self.tok_emb = nn.Embedding(vocab_size, d_model)
97
+ self.pos_emb = nn.Embedding(max_seq_len, d_model)
98
+ self.drop = nn.Dropout(dropout)
99
+ self.blocks = nn.ModuleList(
100
+ [Block(d_model, heads, ff_mult, dropout) for _ in range(layers)]
101
+ )
102
+ self.ln_f = nn.LayerNorm(d_model)
103
+
104
+ if task == "text-generation":
105
+ self.head = nn.Linear(d_model, vocab_size, bias=False)
106
+ self.head.weight = self.tok_emb.weight # weight tying
107
+ elif task == "text-classification":
108
+ assert n_classes > 0, "n_classes must be set for text-classification"
109
+ self.head = nn.Linear(d_model, n_classes)
110
+ else:
111
+ raise ValueError(f"Unsupported task for TinyTransformer: {task}")
112
+
113
+ self.apply(self._init_weights)
114
+
115
+ def _init_weights(self, module: nn.Module) -> None:
116
+ if isinstance(module, nn.Linear):
117
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
118
+ if module.bias is not None:
119
+ nn.init.zeros_(module.bias)
120
+ elif isinstance(module, nn.Embedding):
121
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
122
+
123
+ def forward(self, input_ids: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
124
+ B, T = input_ids.shape
125
+ assert T <= self.max_seq_len, (
126
+ f"Sequence length {T} exceeds max_seq_len {self.max_seq_len}"
127
+ )
128
+ pos = torch.arange(T, device=input_ids.device).unsqueeze(0)
129
+ x = self.tok_emb(input_ids) + self.pos_emb(pos)
130
+ x = self.drop(x)
131
+
132
+ attn_mask = None
133
+ if attention_mask is not None:
134
+ # Combine causal mask with padding mask.
135
+ causal = torch.tril(torch.ones(T, T, device=input_ids.device, dtype=torch.bool))
136
+ pad = attention_mask.bool().unsqueeze(1).unsqueeze(1) # B,1,1,T
137
+ attn_mask = (causal.unsqueeze(0).unsqueeze(0) & pad)
138
+
139
+ for block in self.blocks:
140
+ x = block(x, attn_mask=attn_mask)
141
+ x = self.ln_f(x)
142
+
143
+ if self.task == "text-generation":
144
+ return self.head(x) # B, T, vocab_size
145
+ else:
146
+ if attention_mask is not None:
147
+ lengths = attention_mask.sum(dim=1).clamp(min=1) - 1
148
+ else:
149
+ lengths = torch.full((B,), T - 1, device=input_ids.device)
150
+ pooled = x[torch.arange(B, device=input_ids.device), lengths]
151
+ return self.head(pooled) # B, n_classes
152
+
153
+ @torch.no_grad()
154
+ def generate(
155
+ self,
156
+ input_ids: torch.Tensor,
157
+ max_new_tokens: int,
158
+ temperature: float = 0.8,
159
+ top_k: Optional[int] = 40,
160
+ eos_id: Optional[int] = None,
161
+ ) -> torch.Tensor:
162
+ self.eval()
163
+ for _ in range(max_new_tokens):
164
+ cond = input_ids[:, -self.max_seq_len:]
165
+ logits = self(cond)
166
+ logits = logits[:, -1, :] / max(temperature, 1e-5)
167
+ if top_k is not None:
168
+ v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
169
+ logits[logits < v[:, [-1]]] = -float("inf")
170
+ probs = F.softmax(logits, dim=-1)
171
+ next_id = torch.multinomial(probs, num_samples=1)
172
+ input_ids = torch.cat([input_ids, next_id], dim=1)
173
+ if eos_id is not None and (next_id == eos_id).all():
174
+ break
175
+ return input_ids