mostlyai-engine 1.5.4__tar.gz → 1.5.6__tar.gz
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.
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/PKG-INFO +1 -1
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/__init__.py +1 -1
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_common.py +10 -5
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/argn.py +36 -18
- mostlyai_engine-1.5.6/mostlyai/engine/_tabular/common.py +37 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/generation.py +21 -15
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/training.py +1 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/pyproject.toml +1 -1
- mostlyai_engine-1.5.4/mostlyai/engine/_tabular/common.py +0 -30
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/.gitignore +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/LICENSE +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/README.md +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_dtypes.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/__init__.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/text.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/__init__.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/common.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/encoding.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/__init__.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/base.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/generation.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/lstm.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/training.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_memory.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/__init__.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/encoding.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/fairness.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_training_utils.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_workspace.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/analysis.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/domain.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/encoding.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/generation.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/logging.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/random_state.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/splitting.py +0 -0
- {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/training.py +0 -0
|
@@ -22,7 +22,7 @@ from mostlyai.engine.splitting import split
|
|
|
22
22
|
from mostlyai.engine.training import train
|
|
23
23
|
|
|
24
24
|
__all__ = ["split", "analyze", "encode", "train", "generate", "init_logging", "set_random_state"]
|
|
25
|
-
__version__ = "1.5.
|
|
25
|
+
__version__ = "1.5.6"
|
|
26
26
|
|
|
27
27
|
# suppress specific warning related to os.fork() in multi-threaded processes
|
|
28
28
|
warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*multi-threaded.*fork.*")
|
|
@@ -56,6 +56,11 @@ SLEN_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}slen_" # seque
|
|
|
56
56
|
SDEC_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}sdec_" # sequence index decile
|
|
57
57
|
TABLE_COLUMN_INFIX = "::" # this should be consistent as in mostly-data and mostlyai-qa
|
|
58
58
|
|
|
59
|
+
# the latest version of the model uses SIDX/SLEN/RIDX positional column
|
|
60
|
+
DEFAULT_HAS_SLEN = True
|
|
61
|
+
DEFAULT_HAS_RIDX = True
|
|
62
|
+
DEFAULT_HAS_SDEC = False
|
|
63
|
+
|
|
59
64
|
ANALYZE_MIN_MAX_TOP_N = 1000 # the number of min/max values to be kept from each partition
|
|
60
65
|
|
|
61
66
|
# the minimal number of min/max values to trigger the reduction; if less, the min/max will be reduced to None
|
|
@@ -318,6 +323,7 @@ def get_argn_name(
|
|
|
318
323
|
def get_cardinalities(
|
|
319
324
|
stats: dict, has_slen: bool | None = None, has_ridx: bool | None = None, has_sdec: bool | None = None
|
|
320
325
|
) -> dict[str, int]:
|
|
326
|
+
# the latest version of the model uses SIDX/SLEN/RIDX positional column (applies to sequential model only)
|
|
321
327
|
cardinalities: dict[str, int] = {}
|
|
322
328
|
|
|
323
329
|
if stats.get("is_sequential", False):
|
|
@@ -538,12 +544,11 @@ def decode_positional_column(df_encoded: pd.DataFrame, max_seq_len: int, prefix:
|
|
|
538
544
|
|
|
539
545
|
|
|
540
546
|
def get_positional_cardinalities(
|
|
541
|
-
max_seq_len: int, has_slen: bool | None, has_ridx: bool | None, has_sdec: bool | None
|
|
547
|
+
max_seq_len: int, has_slen: bool | None = None, has_ridx: bool | None = None, has_sdec: bool | None = None
|
|
542
548
|
) -> dict[str, int]:
|
|
543
|
-
|
|
544
|
-
|
|
545
|
-
|
|
546
|
-
has_sdec = has_sdec if has_sdec is not None else False
|
|
549
|
+
has_slen = has_slen if has_slen is not None else DEFAULT_HAS_SLEN
|
|
550
|
+
has_ridx = has_ridx if has_ridx is not None else DEFAULT_HAS_RIDX
|
|
551
|
+
has_sdec = has_sdec if has_sdec is not None else DEFAULT_HAS_SDEC
|
|
547
552
|
|
|
548
553
|
if max_seq_len < SIDX_RIDX_DIGIT_ENCODING_THRESHOLD:
|
|
549
554
|
# encode positional columns as numeric_discrete
|
|
@@ -226,15 +226,12 @@ class Embedders(nn.Module):
|
|
|
226
226
|
),
|
|
227
227
|
None,
|
|
228
228
|
)
|
|
229
|
-
last_ridx_sub_col = next(
|
|
230
|
-
(sub_col for sub_col in reversed(self.cardinalities) if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX)), None
|
|
231
|
-
)
|
|
232
229
|
for sub_col, dim_input in self.cardinalities.items():
|
|
233
230
|
dim_output = _embedding_heuristic(id=self.id(sub_col), model_size=model_size, dim_input=dim_input)
|
|
234
231
|
embedder = nn.Embedding(num_embeddings=dim_input, embedding_dim=dim_output, device=device)
|
|
235
|
-
# the embeddings of the last
|
|
232
|
+
# the embeddings of the last SLEN sub column are never used
|
|
236
233
|
# so we explicitly freeze them to make opacus not complain about "per sample gradient is not initialized"
|
|
237
|
-
if sub_col
|
|
234
|
+
if sub_col == last_slen_sub_col:
|
|
238
235
|
embedder.weight.requires_grad = False
|
|
239
236
|
self.add(sub_column=sub_col, embedder=embedder)
|
|
240
237
|
self.dims.append(dim_output)
|
|
@@ -874,6 +871,7 @@ class FlatModel(nn.Module):
|
|
|
874
871
|
model_size: ModelSizeOrUnits,
|
|
875
872
|
column_order: list[str] | None,
|
|
876
873
|
device: torch.device,
|
|
874
|
+
with_dp: bool = False,
|
|
877
875
|
):
|
|
878
876
|
super().__init__()
|
|
879
877
|
|
|
@@ -895,6 +893,7 @@ class FlatModel(nn.Module):
|
|
|
895
893
|
ctx_cardinalities=self.ctx_cardinalities,
|
|
896
894
|
ctxseq_len_median=self.ctxseq_len_median,
|
|
897
895
|
device=device,
|
|
896
|
+
with_dp=with_dp,
|
|
898
897
|
)
|
|
899
898
|
|
|
900
899
|
# sub column embeddings
|
|
@@ -1256,26 +1255,28 @@ class SequentialModel(nn.Module):
|
|
|
1256
1255
|
if context is None:
|
|
1257
1256
|
context = self.context_compressor(x)
|
|
1258
1257
|
|
|
1259
|
-
# SLEN is only masked for models with RIDX
|
|
1260
1258
|
has_ridx = any(sub_col.startswith(RIDX_SUB_COLUMN_PREFIX) for sub_col in self.tgt_cardinalities)
|
|
1261
|
-
|
|
1259
|
+
|
|
1260
|
+
# SLEN and RIDX are masked for history
|
|
1261
|
+
# NOTE: SLEN is not masked for models without RIDX (backwards compatibility)
|
|
1262
|
+
history_masked_sub_cols = (
|
|
1262
1263
|
(SLEN_SUB_COLUMN_PREFIX, RIDX_SUB_COLUMN_PREFIX) if has_ridx else (RIDX_SUB_COLUMN_PREFIX,)
|
|
1263
1264
|
)
|
|
1265
|
+
# SLEN is masked for column embeddings
|
|
1266
|
+
# NOTE: SLEN is not masked for models without RIDX (backwards compatibility)
|
|
1267
|
+
col_embeddings_masked_sub_cols = (SLEN_SUB_COLUMN_PREFIX,) if has_ridx else ()
|
|
1264
1268
|
|
|
1265
1269
|
outputs = {}
|
|
1266
1270
|
if mode == "trn":
|
|
1267
1271
|
# forward pass through sub column embedders
|
|
1268
1272
|
tgt_embeds = self.embedders(x)
|
|
1269
|
-
tgt_embeds_positional_masked = {
|
|
1270
|
-
k: torch.zeros_like(v) if k.startswith(masked_positional_columns) else v for k, v in tgt_embeds.items()
|
|
1271
|
-
}
|
|
1272
|
-
|
|
1273
|
-
# forward pass through column embedders
|
|
1274
|
-
tgt_col_embeds = self.column_embedders(tgt_embeds_positional_masked)
|
|
1275
1273
|
|
|
1276
1274
|
# history
|
|
1277
1275
|
# time shift: remove last time step; add zeros for first time step; add randoms for all others
|
|
1278
|
-
|
|
1276
|
+
tgt_embeds_for_history = {
|
|
1277
|
+
k: torch.zeros_like(v) if k.startswith(history_masked_sub_cols) else v for k, v in tgt_embeds.items()
|
|
1278
|
+
}
|
|
1279
|
+
embeddings = torch.cat(list(tgt_embeds_for_history.values()), dim=-1)
|
|
1279
1280
|
history_in = embeddings[:, :-1, :]
|
|
1280
1281
|
history_in = nn.ConstantPad2d((0, 0, 1, 0), 0)(history_in)
|
|
1281
1282
|
history, _ = self.history_compressor(history_in)
|
|
@@ -1287,6 +1288,13 @@ class SequentialModel(nn.Module):
|
|
|
1287
1288
|
flat_ctx = self._repeat_flat_context(flat_ctx, history.size(1))
|
|
1288
1289
|
context_history = [torch.cat(flat_ctx + seq_ctx + [history], -1)]
|
|
1289
1290
|
|
|
1291
|
+
# forward pass through column embedders
|
|
1292
|
+
tgt_embeds_for_col_embeddings = {
|
|
1293
|
+
k: torch.zeros_like(v) if k.startswith(col_embeddings_masked_sub_cols) else v
|
|
1294
|
+
for k, v in tgt_embeds.items()
|
|
1295
|
+
}
|
|
1296
|
+
tgt_col_embeds = self.column_embedders(tgt_embeds_for_col_embeddings)
|
|
1297
|
+
|
|
1290
1298
|
# create batch-wise permutation mask
|
|
1291
1299
|
col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
|
|
1292
1300
|
col_mask = _make_permutation_mask(
|
|
@@ -1327,7 +1335,9 @@ class SequentialModel(nn.Module):
|
|
|
1327
1335
|
return outputs, {}
|
|
1328
1336
|
|
|
1329
1337
|
else: # mode == "gen"
|
|
1338
|
+
is_0th_step = False
|
|
1330
1339
|
if history is None or history_state is None:
|
|
1340
|
+
is_0th_step = True
|
|
1331
1341
|
# initialize history
|
|
1332
1342
|
history_in = torch.cat(
|
|
1333
1343
|
[
|
|
@@ -1401,24 +1411,32 @@ class SequentialModel(nn.Module):
|
|
|
1401
1411
|
)
|
|
1402
1412
|
|
|
1403
1413
|
# update timestep output
|
|
1414
|
+
if is_0th_step and sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
|
|
1415
|
+
# overwrite output for RIDX sub-columns on 0th step with SLEN sub-columns
|
|
1416
|
+
slen_sub_col = sub_col.replace(RIDX_SUB_COLUMN_PREFIX, SLEN_SUB_COLUMN_PREFIX)
|
|
1417
|
+
out = outputs[slen_sub_col] if slen_sub_col in outputs else out
|
|
1404
1418
|
outputs[sub_col] = out
|
|
1405
1419
|
|
|
1406
1420
|
# update current sub column embedding
|
|
1407
1421
|
tgt_embeds[sub_col] = self.embedders.get(sub_col)(out)
|
|
1408
|
-
|
|
1409
|
-
k: torch.zeros_like(v) if k.startswith(
|
|
1422
|
+
tgt_embeds_for_history = {
|
|
1423
|
+
k: torch.zeros_like(v) if k.startswith(history_masked_sub_cols) else v
|
|
1424
|
+
for k, v in tgt_embeds.items()
|
|
1425
|
+
}
|
|
1426
|
+
tgt_embeds_for_col_embeddings = {
|
|
1427
|
+
k: torch.zeros_like(v) if k.startswith(col_embeddings_masked_sub_cols) else v
|
|
1410
1428
|
for k, v in tgt_embeds.items()
|
|
1411
1429
|
}
|
|
1412
1430
|
|
|
1413
1431
|
# update current column embedding
|
|
1414
1432
|
if sub_col in self.tgt_last_sub_cols:
|
|
1415
1433
|
col_sub_cols = self.tgt_column_sub_columns[lookup.col_name]
|
|
1416
|
-
col_embed_in = torch.cat([
|
|
1434
|
+
col_embed_in = torch.cat([tgt_embeds_for_col_embeddings[sc] for sc in col_sub_cols], dim=-1)
|
|
1417
1435
|
tgt_col_embeds[lookup.col_name] = self.column_embedders.get(lookup.col_name)(col_embed_in)
|
|
1418
1436
|
col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
|
|
1419
1437
|
|
|
1420
1438
|
# update history and hidden state
|
|
1421
|
-
history_in = torch.cat(
|
|
1439
|
+
history_in = torch.cat(list(tgt_embeds_for_history.values()), dim=-1)
|
|
1422
1440
|
history, history_state = self.history_compressor(history_in, history_state=history_state)
|
|
1423
1441
|
|
|
1424
1442
|
# order outputs according to tgt_sub_columns
|
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
# Copyright 2025 MOSTLY AI
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
import logging
|
|
16
|
+
import time
|
|
17
|
+
from pathlib import Path
|
|
18
|
+
|
|
19
|
+
import torch
|
|
20
|
+
|
|
21
|
+
_LOG = logging.getLogger(__name__)
|
|
22
|
+
|
|
23
|
+
DPLSTM_SUFFIXES: tuple = ("ih.weight", "ih.bias", "hh.weight", "hh.bias")
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def load_model_weights(model: torch.nn.Module, path: Path, device: torch.device) -> None:
|
|
27
|
+
t0 = time.time()
|
|
28
|
+
incompatible_keys = model.load_state_dict(torch.load(f=path, map_location=device, weights_only=True), strict=False)
|
|
29
|
+
missing_keys = incompatible_keys.missing_keys
|
|
30
|
+
unexpected_keys = incompatible_keys.unexpected_keys
|
|
31
|
+
# for DP-trained models, we expect extra keys from the DPLSTM layers (which is fine to ignore because we use standard LSTM layers during generation)
|
|
32
|
+
# but if there're any other missing or unexpected keys, an error should be raised
|
|
33
|
+
if len(missing_keys) > 0 or any(not k.endswith(DPLSTM_SUFFIXES) for k in unexpected_keys):
|
|
34
|
+
raise RuntimeError(
|
|
35
|
+
f"failed to load model weights due to incompatibility: {missing_keys = }, {unexpected_keys = }"
|
|
36
|
+
)
|
|
37
|
+
_LOG.info(f"loaded model weights in {time.time() - t0:.2f}s")
|
|
@@ -30,6 +30,9 @@ from mostlyai.engine._common import (
|
|
|
30
30
|
ARGN_TABLE,
|
|
31
31
|
CTXFLT,
|
|
32
32
|
CTXSEQ,
|
|
33
|
+
DEFAULT_HAS_RIDX,
|
|
34
|
+
DEFAULT_HAS_SDEC,
|
|
35
|
+
DEFAULT_HAS_SLEN,
|
|
33
36
|
POSITIONAL_COLUMN,
|
|
34
37
|
RIDX_SUB_COLUMN_PREFIX,
|
|
35
38
|
SDEC_SUB_COLUMN_PREFIX,
|
|
@@ -695,6 +698,8 @@ def generate(
|
|
|
695
698
|
has_slen = any(SLEN_SUB_COLUMN_PREFIX in k for k in model_units.keys())
|
|
696
699
|
has_ridx = any(RIDX_SUB_COLUMN_PREFIX in k for k in model_units.keys())
|
|
697
700
|
has_sdec = any(SDEC_SUB_COLUMN_PREFIX in k for k in model_units.keys())
|
|
701
|
+
else:
|
|
702
|
+
has_slen, has_ridx, has_sdec = DEFAULT_HAS_SLEN, DEFAULT_HAS_RIDX, DEFAULT_HAS_SDEC
|
|
698
703
|
|
|
699
704
|
tgt_cardinalities = get_cardinalities(tgt_stats, has_slen, has_ridx, has_sdec)
|
|
700
705
|
ctx_cardinalities = get_cardinalities(ctx_stats)
|
|
@@ -874,11 +879,14 @@ def generate(
|
|
|
874
879
|
no_of_model_params = get_no_of_model_parameters(model)
|
|
875
880
|
_LOG.info(f"{no_of_model_params=}")
|
|
876
881
|
|
|
877
|
-
|
|
878
|
-
|
|
879
|
-
|
|
880
|
-
|
|
881
|
-
|
|
882
|
+
if workspace.model_tabular_weights_path.exists():
|
|
883
|
+
load_model_weights(
|
|
884
|
+
model=model,
|
|
885
|
+
path=workspace.model_tabular_weights_path,
|
|
886
|
+
device=device,
|
|
887
|
+
)
|
|
888
|
+
else:
|
|
889
|
+
_LOG.warning("model weights not found; generating data with an untrained model")
|
|
882
890
|
|
|
883
891
|
model.to(device)
|
|
884
892
|
model.eval()
|
|
@@ -1023,7 +1031,7 @@ def generate(
|
|
|
1023
1031
|
|
|
1024
1032
|
# fix SLEN by propagating sampled SLEN from first step
|
|
1025
1033
|
slen_vals = {}
|
|
1026
|
-
if seq_step > 0:
|
|
1034
|
+
if has_slen and seq_step > 0:
|
|
1027
1035
|
slen = out_df[SLEN_SUB_COLUMN_PREFIX]
|
|
1028
1036
|
slen = encode_positional_column(slen, max_seq_len=seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX)
|
|
1029
1037
|
slen_vals = {
|
|
@@ -1102,16 +1110,14 @@ def generate(
|
|
|
1102
1110
|
out_df[SIDX_SUB_COLUMN_PREFIX] = decode_positional_column(
|
|
1103
1111
|
out_df, seq_len_max, prefix=SIDX_SUB_COLUMN_PREFIX
|
|
1104
1112
|
)
|
|
1105
|
-
|
|
1106
|
-
out_df
|
|
1107
|
-
|
|
1108
|
-
|
|
1113
|
+
if has_slen:
|
|
1114
|
+
out_df[SLEN_SUB_COLUMN_PREFIX] = decode_positional_column(
|
|
1115
|
+
out_df, seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX
|
|
1116
|
+
)
|
|
1117
|
+
out_df[SLEN_SUB_COLUMN_PREFIX] = out_df[SLEN_SUB_COLUMN_PREFIX].clip(lower=seq_len_min)
|
|
1109
1118
|
if has_ridx:
|
|
1110
|
-
|
|
1111
|
-
|
|
1112
|
-
decode_positional_column(out_df, seq_len_max, prefix=RIDX_SUB_COLUMN_PREFIX)
|
|
1113
|
-
if seq_step > 0
|
|
1114
|
-
else out_df[SLEN_SUB_COLUMN_PREFIX]
|
|
1119
|
+
out_df[RIDX_SUB_COLUMN_PREFIX] = decode_positional_column(
|
|
1120
|
+
out_df, seq_len_max, prefix=RIDX_SUB_COLUMN_PREFIX
|
|
1115
1121
|
)
|
|
1116
1122
|
out_df[RIDX_SUB_COLUMN_PREFIX] = out_df[RIDX_SUB_COLUMN_PREFIX].clip(
|
|
1117
1123
|
lower=seq_len_min - seq_step, upper=seq_len_max
|
|
@@ -1,30 +0,0 @@
|
|
|
1
|
-
# Copyright 2025 MOSTLY AI
|
|
2
|
-
#
|
|
3
|
-
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
-
# you may not use this file except in compliance with the License.
|
|
5
|
-
# You may obtain a copy of the License at
|
|
6
|
-
#
|
|
7
|
-
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
-
#
|
|
9
|
-
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
-
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
-
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
-
# See the License for the specific language governing permissions and
|
|
13
|
-
# limitations under the License.
|
|
14
|
-
|
|
15
|
-
import logging
|
|
16
|
-
import time
|
|
17
|
-
from pathlib import Path
|
|
18
|
-
|
|
19
|
-
import torch
|
|
20
|
-
|
|
21
|
-
_LOG = logging.getLogger(__name__)
|
|
22
|
-
|
|
23
|
-
|
|
24
|
-
def load_model_weights(model: torch.nn.Module, path: Path, device: torch.device):
|
|
25
|
-
try:
|
|
26
|
-
t00 = time.time()
|
|
27
|
-
model.load_state_dict(torch.load(f=path, map_location=device, weights_only=True))
|
|
28
|
-
_LOG.info(f"loaded model weights in {time.time() - t00:.2f}s")
|
|
29
|
-
except Exception as e:
|
|
30
|
-
_LOG.warning(f"failed to load model weights: {e}")
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/numeric.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/text.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/character.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/itt.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/lat_long.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/numeric.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/hf_engine.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/vllm_engine.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/tokenizer_utils.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|