tensorless 0.6.0__tar.gz → 0.7.0__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.
Files changed (84) hide show
  1. {tensorless-0.6.0/tensorless.egg-info → tensorless-0.7.0}/PKG-INFO +1 -2
  2. {tensorless-0.6.0 → tensorless-0.7.0}/README.md +1 -2
  3. {tensorless-0.6.0 → tensorless-0.7.0}/docs/training.md +22 -0
  4. {tensorless-0.6.0 → tensorless-0.7.0}/pyproject.toml +1 -1
  5. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/__init__.py +2 -1
  6. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/api.py +26 -1
  7. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/auto/config.py +1 -0
  8. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/config.py +2 -0
  9. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/engine.py +6 -3
  10. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/serialization/tl_format.py +1 -0
  11. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/training/trainer.py +13 -3
  12. {tensorless-0.6.0 → tensorless-0.7.0/tensorless.egg-info}/PKG-INFO +1 -2
  13. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_train_text_generation.py +15 -0
  14. {tensorless-0.6.0 → tensorless-0.7.0}/LICENSE +0 -0
  15. {tensorless-0.6.0 → tensorless-0.7.0}/MANIFEST.in +0 -0
  16. {tensorless-0.6.0 → tensorless-0.7.0}/docs/api_reference.md +0 -0
  17. {tensorless-0.6.0 → tensorless-0.7.0}/docs/architecture.md +0 -0
  18. {tensorless-0.6.0 → tensorless-0.7.0}/docs/automatic_mode.md +0 -0
  19. {tensorless-0.6.0 → tensorless-0.7.0}/docs/checkpointing.md +0 -0
  20. {tensorless-0.6.0 → tensorless-0.7.0}/docs/cli.md +0 -0
  21. {tensorless-0.6.0 → tensorless-0.7.0}/docs/configuration.md +0 -0
  22. {tensorless-0.6.0 → tensorless-0.7.0}/docs/contributing.md +0 -0
  23. {tensorless-0.6.0 → tensorless-0.7.0}/docs/examples.md +0 -0
  24. {tensorless-0.6.0 → tensorless-0.7.0}/docs/inference.md +0 -0
  25. {tensorless-0.6.0 → tensorless-0.7.0}/docs/installation.md +0 -0
  26. {tensorless-0.6.0 → tensorless-0.7.0}/docs/limitations.md +0 -0
  27. {tensorless-0.6.0 → tensorless-0.7.0}/docs/quickstart.md +0 -0
  28. {tensorless-0.6.0 → tensorless-0.7.0}/docs/roadmap.md +0 -0
  29. {tensorless-0.6.0 → tensorless-0.7.0}/docs/tl_format.md +0 -0
  30. {tensorless-0.6.0 → tensorless-0.7.0}/docs/troubleshooting.md +0 -0
  31. {tensorless-0.6.0 → tensorless-0.7.0}/docs/tutorial.md +0 -0
  32. {tensorless-0.6.0 → tensorless-0.7.0}/examples/tabular_classification_example.py +0 -0
  33. {tensorless-0.6.0 → tensorless-0.7.0}/examples/tabular_regression_example.py +0 -0
  34. {tensorless-0.6.0 → tensorless-0.7.0}/examples/text_classification_example.py +0 -0
  35. {tensorless-0.6.0 → tensorless-0.7.0}/examples/text_generation_example.py +0 -0
  36. {tensorless-0.6.0 → tensorless-0.7.0}/setup.cfg +0 -0
  37. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/_version.py +0 -0
  38. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/auto/__init__.py +0 -0
  39. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/auto/detector.py +0 -0
  40. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/backends/__init__.py +0 -0
  41. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/backends/jax_backend.py +0 -0
  42. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/backends/mlx_backend.py +0 -0
  43. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/checkpoint/__init__.py +0 -0
  44. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/checkpoint/manager.py +0 -0
  45. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/cli/__init__.py +0 -0
  46. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/cli/main.py +0 -0
  47. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/data/__init__.py +0 -0
  48. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/data/english_grammar.txt +0 -0
  49. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/data/fingerprint.py +0 -0
  50. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/data/inspector.py +0 -0
  51. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/data/loader.py +0 -0
  52. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/data/tabular.py +0 -0
  53. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/devices/__init__.py +0 -0
  54. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/devices/device.py +0 -0
  55. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/devices/memory.py +0 -0
  56. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/errors.py +0 -0
  57. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/models/__init__.py +0 -0
  58. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/models/mlp.py +0 -0
  59. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/models/registry.py +0 -0
  60. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/models/transformer.py +0 -0
  61. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/runtime.py +0 -0
  62. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/serialization/__init__.py +0 -0
  63. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/tokenization/__init__.py +0 -0
  64. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/tokenization/bpe_tokenizer.py +0 -0
  65. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/tokenization/char_tokenizer.py +0 -0
  66. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/training/__init__.py +0 -0
  67. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/training/data_prep.py +0 -0
  68. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless/training/early_stopping.py +0 -0
  69. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless.egg-info/SOURCES.txt +0 -0
  70. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless.egg-info/dependency_links.txt +0 -0
  71. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless.egg-info/entry_points.txt +0 -0
  72. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless.egg-info/requires.txt +0 -0
  73. {tensorless-0.6.0 → tensorless-0.7.0}/tensorless.egg-info/top_level.txt +0 -0
  74. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_auto_detection.py +0 -0
  75. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_checkpoint_resume.py +0 -0
  76. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_cli.py +0 -0
  77. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_data_loading.py +0 -0
  78. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_end_to_end.py +0 -0
  79. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_fingerprint.py +0 -0
  80. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_jax_backend.py +0 -0
  81. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_mlx_backend.py +0 -0
  82. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_serialization.py +0 -0
  83. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_train_tabular.py +0 -0
  84. {tensorless-0.6.0 → tensorless-0.7.0}/tests/test_train_text_classification.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tensorless
3
- Version: 0.6.0
3
+ Version: 0.7.0
4
4
  Summary: ML with maximum automation and minimum setup.
5
5
  Author: Tensorless Contributors
6
6
  License: MIT
@@ -105,4 +105,3 @@ without a PyTorch dependency. Accelerator cache helpers are available as
105
105
 
106
106
  See the [documentation](docs/quickstart.md) for data formats, configuration,
107
107
  checkpointing, and the command-line interface.
108
- pypi-AgEIcHlwaS5vcmcCJDJmYzZiYmQ1LTI1YTAtNGNlZi05OWE2LTliMjg3ZjY5MThiZQACElsxLFsidGVuc29ybGVzcyJdXQACLFsyLFsiNjcxZGZmNDQtZmVmMC00MWNiLWIzYzYtZWQzMTI3OWU4NWM5Il1dAAAGIDZAp3AAlp0idrTOMPJ227qF_7W0LDsyUpi6pfqf0aV0
@@ -78,5 +78,4 @@ without a PyTorch dependency. Accelerator cache helpers are available as
78
78
  `tensorless.devices.clear_memory()` and `tensorless.devices.memory_stats()`.
79
79
 
80
80
  See the [documentation](docs/quickstart.md) for data formats, configuration,
81
- checkpointing, and the command-line interface.
82
- pypi-AgEIcHlwaS5vcmcCJDJmYzZiYmQ1LTI1YTAtNGNlZi05OWE2LTliMjg3ZjY5MThiZQACElsxLFsidGVuc29ybGVzcyJdXQACLFsyLFsiNjcxZGZmNDQtZmVmMC00MWNiLWIzYzYtZWQzMTI3OWU4NWM5Il1dAAAGIDZAp3AAlp0idrTOMPJ227qF_7W0LDsyUpi6pfqf0aV0
81
+ checkpointing, and the command-line interface.
@@ -12,6 +12,28 @@ Returns a `LoadedModel` (see [inference.md](inference.md)) ready for
12
12
  predictions, and writes `model.tl` (plus a `model.tl.ckpt/` checkpoint
13
13
  directory) to the current directory.
14
14
 
15
+ ## Pretraining and fine-tuning
16
+
17
+ Use the built-in corpus (or any text corpus) to create a base model, then
18
+ continue training its learned weights on your own data:
19
+
20
+ ```python
21
+ base = tl.pretrain(out="english_pretrained.tl", epochs=10)
22
+ model = tl.train(
23
+ "./my_text.txt",
24
+ pretrained="english_pretrained.tl",
25
+ out="my_model.tl",
26
+ epochs=5,
27
+ learning_rate=1e-4,
28
+ )
29
+ ```
30
+
31
+ The source tokenizer and compatible architecture are reused automatically.
32
+ The fine-tuning run starts a fresh optimizer and scheduler; interrupted
33
+ fine-tuning still resumes from its own checkpoint. Use
34
+ `tl.load_pretrained("english_pretrained.tl")` when you want to load the base
35
+ model directly for inference.
36
+
15
37
  ## Supported data formats
16
38
 
17
39
  | Format | Notes |
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "tensorless"
7
- version = "0.6.0"
7
+ version = "0.7.0"
8
8
  description = "ML with maximum automation and minimum setup."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.9"
@@ -12,7 +12,7 @@ ML with maximum automation and minimum setup.
12
12
  See https://github.com/tensorless/tensorless for full documentation.
13
13
  """
14
14
 
15
- from .api import train, pretrain, run, load, inspect
15
+ from .api import train, pretrain, run, load, load_pretrained, inspect
16
16
  from .config import TrainConfig
17
17
  from .errors import (
18
18
  TensorlessError,
@@ -29,6 +29,7 @@ __all__ = [
29
29
  "pretrain",
30
30
  "run",
31
31
  "load",
32
+ "load_pretrained",
32
33
  "inspect",
33
34
  "TrainConfig",
34
35
  "TensorlessError",
@@ -55,7 +55,8 @@ def inspect(path: str) -> InspectionReport:
55
55
  def train(path: str, **kwargs: Any) -> LoadedModel:
56
56
  """Train a model on the dataset at `path`, fully automatically by
57
57
  default. Any field of `TrainConfig` can be overridden via keyword
58
- argument, e.g. `tl.train("./data", d_model=512, layers=6)`.
58
+ argument, e.g. `tl.train("./data", d_model=512, layers=6)`. Pass
59
+ `pretrained="base.tl"` to fine-tune an existing Tensorless model.
59
60
 
60
61
  Implements the "Smart Auto Check":
61
62
  - if an up-to-date trained model already exists for this exact
@@ -65,6 +66,14 @@ def train(path: str, **kwargs: Any) -> LoadedModel:
65
66
  `ask_on_data_change=True`), unless `force=True` is passed
66
67
  """
67
68
  user_cfg = _build_train_config(**kwargs)
69
+ pretrained_state = None
70
+ if user_cfg.pretrained:
71
+ try:
72
+ pretrained_state = load_tl(user_cfg.pretrained)
73
+ except Exception as exc:
74
+ raise ModelError(
75
+ f"Could not load pretrained model '{user_cfg.pretrained}': {exc}"
76
+ ) from exc
68
77
  out = user_cfg.out or "model.tl"
69
78
  checkpoint_dir = user_cfg.checkpoint_dir or (out + ".ckpt")
70
79
  checkpoint_mgr = CheckpointManager(checkpoint_dir)
@@ -130,6 +139,16 @@ def train(path: str, **kwargs: Any) -> LoadedModel:
130
139
  ds = load_dataset(path)
131
140
  resolved = resolve_config(ds, user_cfg)
132
141
  cfg = resolved.to_dict()
142
+ if pretrained_state is not None:
143
+ source_cfg = pretrained_state["config"]
144
+ for field in ("d_model", "layers", "heads", "ff_mult", "max_seq_len", "tokenizer", "bpe_vocab_size"):
145
+ if field not in user_cfg.overrides() and field in source_cfg:
146
+ cfg[field] = source_cfg[field]
147
+ if pretrained_state["task"] != cfg["task"] or pretrained_state["model_type"] != cfg["model_type"]:
148
+ raise ModelError(
149
+ "Pretrained model and target training data must use the same "
150
+ "task and model_type."
151
+ )
133
152
 
134
153
  if resume_state is not None:
135
154
  # Resumed runs must keep the exact architecture/config used
@@ -143,6 +162,7 @@ def train(path: str, **kwargs: Any) -> LoadedModel:
143
162
  checkpoint_mgr=checkpoint_mgr,
144
163
  dataset_fingerprint=fingerprint,
145
164
  resume_state=resume_state,
165
+ pretrained_state=pretrained_state if resume_state is None else None,
146
166
  log_fn=print if cfg.get("verbose", True) else (lambda *a, **k: None),
147
167
  )
148
168
 
@@ -205,6 +225,11 @@ def load(path: str, device: Optional[str] = None) -> LoadedModel:
205
225
  return load_model(path, device=device)
206
226
 
207
227
 
228
+ def load_pretrained(path: str, device: Optional[str] = None) -> LoadedModel:
229
+ """Load a portable model intended to be used as a fine-tuning base."""
230
+ return load_model(path, device=device)
231
+
232
+
208
233
  def run(path: str, prompt: Optional[str] = None) -> Any:
209
234
  """Run a trained `.tl` model.
210
235
 
@@ -101,6 +101,7 @@ def resolve_config(ds: Dataset, user: TrainConfig) -> ResolvedConfig:
101
101
  max_seq_len = user.max_seq_len or (256 if ds.kind in ("text", "text_labeled") else 1)
102
102
  resolved = ResolvedConfig(
103
103
  out=out,
104
+ pretrained=user.pretrained,
104
105
  force=bool(user.force),
105
106
  resume=user.resume,
106
107
  ask_on_data_change=bool(user.ask_on_data_change),
@@ -20,6 +20,7 @@ class TrainConfig:
20
20
 
21
21
  # --- output / lifecycle ---
22
22
  out: Optional[str] = None # output .tl path, default "model.tl"
23
+ pretrained: Optional[str] = None # optional .tl weights to fine-tune
23
24
  force: bool = False # force retraining even if unchanged
24
25
  resume: Optional[bool] = None # force/forbid resume (None = auto)
25
26
  ask_on_data_change: bool = False # raise instead of auto-retrain on data change
@@ -81,6 +82,7 @@ class ResolvedConfig:
81
82
  """
82
83
 
83
84
  out: str
85
+ pretrained: Optional[str]
84
86
  force: bool
85
87
  resume: Optional[bool]
86
88
  ask_on_data_change: bool
@@ -3,6 +3,7 @@
3
3
  from __future__ import annotations
4
4
 
5
5
  from typing import Dict, Iterable, Iterator
6
+ import math
6
7
  import numpy as np
7
8
 
8
9
 
@@ -226,10 +227,12 @@ class LambdaScheduler:
226
227
  def _factor(self, step):
227
228
  if self.warmup_steps and step < self.warmup_steps:
228
229
  return (step + 1) / self.warmup_steps
229
- return max(
230
- 0.1,
231
- 1.0 - (step - self.warmup_steps) / max(1, self.total_steps - self.warmup_steps),
230
+ progress = min(
231
+ 1.0,
232
+ max(0.0, (step - self.warmup_steps) /
233
+ max(1, self.total_steps - self.warmup_steps)),
232
234
  )
235
+ return 0.1 + 0.9 * 0.5 * (1.0 + math.cos(math.pi * progress))
233
236
 
234
237
  def step(self):
235
238
  self.step_count += 1
@@ -56,6 +56,7 @@ def _migrate_payload(payload: Dict[str, Any]) -> Dict[str, Any]:
56
56
 
57
57
  config = migrated.get("config")
58
58
  if isinstance(config, dict):
59
+ config.setdefault("pretrained", None)
59
60
  config.setdefault("tokenizer", "char")
60
61
  config.setdefault("bpe_vocab_size", 1000)
61
62
  config.setdefault("precision", "fp32")
@@ -53,13 +53,15 @@ def _show_progress(epoch, epochs, step, total_steps, loss):
53
53
 
54
54
  def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointManager,
55
55
  dataset_fingerprint: str, resume_state: Optional[Dict[str, Any]] = None,
56
+ pretrained_state: Optional[Dict[str, Any]] = None,
56
57
  log_fn=print) -> Dict[str, Any]:
57
58
  np.random.seed(cfg["seed"])
58
59
  task, model_type = cfg["task"], cfg["model_type"]
59
60
  device = get_device(cfg["device"])
60
61
  if cfg["verbose"]: log_fn(f"[tensorless] task={task} model={model_type} device={cfg['device']} precision={cfg['precision']}")
61
- tokenizer = tokenizer_from_state_dict(resume_state["tokenizer_state"]) if resume_state and resume_state.get("tokenizer_state") else None
62
- preprocessor = TabularPreprocessor.from_state_dict(resume_state["preprocessor_state"]) if resume_state and resume_state.get("preprocessor_state") else None
62
+ source_state = resume_state or pretrained_state
63
+ tokenizer = tokenizer_from_state_dict(source_state["tokenizer_state"]) if source_state and source_state.get("tokenizer_state") else None
64
+ preprocessor = TabularPreprocessor.from_state_dict(source_state["preprocessor_state"]) if source_state and source_state.get("preprocessor_state") else None
63
65
  if task == "text-generation": prepared = dp.prepare_text_generation(ds, cfg, tokenizer=tokenizer)
64
66
  elif task == "text-classification": prepared = dp.prepare_text_classification(ds, cfg, tokenizer=tokenizer, classes=resume_state["meta"]["classes"] if resume_state else None)
65
67
  elif task in ("classification", "regression"): prepared = dp.prepare_tabular(ds, cfg, task, preprocessor)
@@ -71,6 +73,14 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
71
73
  elif device == "mps":
72
74
  backend = "mlx"
73
75
  model = build_model(task, model_type, cfg, prepared.meta, backend=backend)
76
+ if pretrained_state is not None:
77
+ try:
78
+ model.load_state_dict(pretrained_state["model_state_dict"])
79
+ except (KeyError, TypeError, ValueError) as exc:
80
+ raise ValueError(
81
+ "Pretrained weights are incompatible with the target model. "
82
+ "Keep the task, tokenizer, and architecture dimensions compatible."
83
+ ) from exc
74
84
  optimizer = _build_optimizer(model, cfg)
75
85
  total_steps = cfg.get("max_steps") or max(1, len(prepared.train_loader)) * cfg["epochs"]
76
86
  scheduler = LambdaScheduler(optimizer, cfg["warmup_steps"], total_steps)
@@ -99,7 +109,7 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
99
109
  if not np.isfinite(last_train_loss):
100
110
  raise FloatingPointError(f"Non-finite training loss at step {global_step + 1}: {last_train_loss}")
101
111
  if cfg["grad_clip"]:
102
- norm = np.sqrt(sum(float(np.sum(p.grad * p.grad)) for p in model.parameters()))
112
+ norm = np.sqrt(sum(float(np.sum(p.grad.astype(np.float64) ** 2)) for p in model.parameters()))
103
113
  if not np.isfinite(norm):
104
114
  raise FloatingPointError(f"Non-finite gradient norm at step {global_step + 1}: {norm}")
105
115
  if norm > cfg["grad_clip"]:
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tensorless
3
- Version: 0.6.0
3
+ Version: 0.7.0
4
4
  Summary: ML with maximum automation and minimum setup.
5
5
  Author: Tensorless Contributors
6
6
  License: MIT
@@ -105,4 +105,3 @@ without a PyTorch dependency. Accelerator cache helpers are available as
105
105
 
106
106
  See the [documentation](docs/quickstart.md) for data formats, configuration,
107
107
  checkpointing, and the command-line interface.
108
- pypi-AgEIcHlwaS5vcmcCJDJmYzZiYmQ1LTI1YTAtNGNlZi05OWE2LTliMjg3ZjY5MThiZQACElsxLFsidGVuc29ybGVzcyJdXQACLFsyLFsiNjcxZGZmNDQtZmVmMC00MWNiLWIzYzYtZWQzMTI3OWU4NWM5Il1dAAAGIDZAp3AAlp0idrTOMPJ227qF_7W0LDsyUpi6pfqf0aV0
@@ -59,6 +59,21 @@ def test_builtin_english_pretraining(workdir):
59
59
  assert model.tokenizer is not None
60
60
 
61
61
 
62
+ def test_finetuning_uses_pretrained_weights(text_corpus, workdir):
63
+ tl.pretrain(
64
+ out="base.tl", epochs=1, max_steps=1, max_seq_len=32,
65
+ d_model=16, layers=1, heads=2, batch_size=2, checkpoint_every=1,
66
+ verbose=False,
67
+ )
68
+ model = tl.train(
69
+ text_corpus, pretrained="base.tl", out="finetuned.tl",
70
+ epochs=1, max_steps=1, batch_size=2, checkpoint_every=1,
71
+ verbose=False,
72
+ )
73
+ assert model.config["d_model"] == 16
74
+ assert tl.load_pretrained("base.tl").task == "text-generation"
75
+
76
+
62
77
  def test_long_text_generation_is_streamed_in_batches():
63
78
  ds = Dataset(kind="text", source="memory", texts=["the quick brown fox " * 2000])
64
79
  cfg = resolve_config(ds, TrainConfig(max_seq_len=32, tokenizer="char"))
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