diff --git a/CHANGELOG.md b/CHANGELOG.md index 2c7218f..161b992 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,6 +10,7 @@ All notable changes to this project will be documented in this file. * Training pipeline with configurable runners * Feature sharding, ddp, sfdp implemented * Compression implemented in activation store (not yet tested properly on downstream effects) +* Support raw text datasets in `ActivationsStore` via lazy tokenization. ### Notes diff --git a/README.md b/README.md index aeed251..b99a8cd 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,10 @@ model = load_model("meta-llama/Llama-3.2-1B", device="cuda") # Create config cfg = clt_training_runner_config() +# To use raw-text datasets instead of pre-tokenized input: +# cfg.is_dataset_tokenized = False +# cfg.dataset_text_column = "text" + # Create activation store store = ActivationsStore(model, cfg) diff --git a/src/clt_forge/config/autointerp_config.py b/src/clt_forge/config/autointerp_config.py index 430a296..b45a5a8 100644 --- a/src/clt_forge/config/autointerp_config.py +++ b/src/clt_forge/config/autointerp_config.py @@ -19,6 +19,7 @@ class AutoInterpConfig(BaseModel): # ---- Dataset ---- dataset_path: str = "" # HuggingFace path or local path is_dataset_tokenized: bool = True + dataset_text_column: str = "text" split: str = "train" # dataset split passed to load_dataset disk: bool = False # use load_from_disk instead of load_dataset is_multilingual_split_dataset: bool = False # only for multilingual setups diff --git a/src/clt_forge/config/clt_training_runner_config.py b/src/clt_forge/config/clt_training_runner_config.py index c80b5c0..148ee36 100644 --- a/src/clt_forge/config/clt_training_runner_config.py +++ b/src/clt_forge/config/clt_training_runner_config.py @@ -24,6 +24,7 @@ class CLTTrainingRunnerConfig(BaseModel): model_from_pretrained_kwargs: Optional[Dict[str, Any]] = None dataset_path: str = "" # Hugging face path is_dataset_tokenized: bool = True + dataset_text_column: str = "text" is_multilingual_split_dataset: bool = False # can be ignored, it is only for multilingual datasets processing split: str = "train" disk: bool = False # use load_from_disk instead and local dataset diff --git a/src/clt_forge/training/activations_store.py b/src/clt_forge/training/activations_store.py index bb3e751..d5079b3 100644 --- a/src/clt_forge/training/activations_store.py +++ b/src/clt_forge/training/activations_store.py @@ -94,21 +94,7 @@ def __init__(self, self.raw_ds = load_dataset_auto(cfg.dataset_path, split=cfg.split, disk=cfg.disk) logger.info("Loaded dataset") - if "tokens" not in self.raw_ds.column_names: - if "input_ids" in self.raw_ds.column_names: - logger.info("tokens column not found — using input_ids instead.") - self.raw_ds = self.raw_ds.rename_column("input_ids", "tokens") - else: - raise ValueError( - f"Dataset {cfg.dataset_path} must contain a pre-tokenised tokens or input_ids column." - ) - - first_tok = self.raw_ds[0]["tokens"] - - if isinstance(first_tok, torch.Tensor) and first_tok.ndim != 1: - raise ValueError("Each 'tokens' entry must be a 1‑D tensor.") - if isinstance(first_tok, (list, tuple)) and any(isinstance(x, list) for x in first_tok): - raise ValueError("Nested sequences detected; expected a flat list of ints.") + self._validate_dataset_columns() self._reset_token_iterator() else: @@ -148,6 +134,96 @@ def __init__(self, assert self.cfg.train_batch_size_tokens % self.cfg.context_size == 0, "ctx size must divide train_batch_size_tokens" # ─────────────────── token pipeline ─────────────────── + def _validate_dataset_columns(self) -> None: + if self.cfg.is_dataset_tokenized: + if "tokens" not in self.raw_ds.column_names: + if "input_ids" in self.raw_ds.column_names: + logger.info("tokens column not found — using input_ids instead.") + self.raw_ds = self.raw_ds.rename_column("input_ids", "tokens") + else: + raise ValueError( + f"Dataset {self.cfg.dataset_path} must contain a pre-tokenised tokens or input_ids column " + "when is_dataset_tokenized=True." + ) + + self._tokens_from_dataset_row(self.raw_ds[0]) + return + + if self.cfg.dataset_text_column not in self.raw_ds.column_names: + raise ValueError( + f"Dataset {self.cfg.dataset_path} must contain dataset_text_column=" + f"'{self.cfg.dataset_text_column}' when is_dataset_tokenized=False." + ) + + if getattr(self.model, "tokenizer", None) is None: + raise ValueError("Raw text datasets require model.tokenizer to tokenize examples.") + + first_text = self.raw_ds[0][self.cfg.dataset_text_column] + if not isinstance(first_text, str): + raise ValueError( + f"Expected dataset_text_column='{self.cfg.dataset_text_column}' to contain strings, " + f"got {type(first_text).__name__}." + ) + + def _tokens_from_dataset_row(self, row: Any) -> torch.Tensor: + if self.cfg.is_dataset_tokenized: + toks = row["tokens"] + else: + toks = self._tokenize_raw_text(row[self.cfg.dataset_text_column]) + + toks = self._ensure_1d_token_tensor(toks) + return self._strip_leading_bos(toks) + + def _tokenize_raw_text(self, text: str) -> Any: + if not isinstance(text, str): + raise ValueError( + f"Expected dataset_text_column='{self.cfg.dataset_text_column}' to contain strings, " + f"got {type(text).__name__}." + ) + + tokenizer = getattr(self.model, "tokenizer", None) + if tokenizer is None: + raise ValueError("Raw text datasets require model.tokenizer to tokenize examples.") + + try: + encoded = tokenizer(text, add_special_tokens=False) + if isinstance(encoded, dict): + return encoded["input_ids"] + if isinstance(encoded, (list, tuple, torch.Tensor)): + return encoded + if hasattr(encoded, "input_ids"): + return encoded.input_ids + except TypeError: + return tokenizer.encode(text, add_special_tokens=False) + + raise ValueError("Tokenizer output must contain input_ids.") + + def _ensure_1d_token_tensor(self, toks: Any) -> torch.Tensor: + if isinstance(toks, torch.Tensor): + toks = toks.detach().cpu().to(dtype=torch.long) + else: + toks = torch.as_tensor(toks, dtype=torch.long) + + if toks.ndim != 1: + raise ValueError("Each token entry must be a 1-D sequence of token ids.") + + return toks + + def _strip_leading_bos(self, toks: torch.Tensor) -> torch.Tensor: + tokenizer = getattr(self.model, "tokenizer", None) + bos_id = None if tokenizer is None else getattr(tokenizer, "bos_token_id", None) + + if bos_id is not None and len(toks) > 0 and toks[0].item() == bos_id: + return toks[1:] + + return toks + + def _truncate_to_context(self, toks: torch.Tensor) -> torch.Tensor: + truncated_len = (len(toks) // self.context_size) * self.context_size + if truncated_len == 0: + return toks + return toks[:truncated_len] + def _iterate_raw_dataset_tokens(self) -> Iterator[torch.Tensor]: """ Yield each row's token vector as a 1‑D torch.Tensor on **CPU**. @@ -163,24 +239,13 @@ def _iterate_raw_dataset_tokens(self) -> Iterator[torch.Tensor]: self.runtime_doc_languages = [] for i in range(start, end): - toks = self.raw_ds[i]["tokens"] - tokenizer = getattr(self.model, "tokenizer", None) - bos_id = None if tokenizer is None else tokenizer.bos_token_id - - if not isinstance(toks, torch.Tensor): - toks = torch.tensor(toks, dtype=torch.long) - - if bos_id is not None and len(toks) > 0 and toks[0].item() == bos_id: - toks = toks[1:] - - doc_len = len(toks) - truncated_len = (doc_len // (self.context_size)) * (self.context_size) # TODO before -1 to both + toks = self._tokens_from_dataset_row(self.raw_ds[i]) + toks = self._truncate_to_context(toks) - if truncated_len > 0: - toks = toks[:truncated_len] + if len(toks) > 0: if self.cfg.is_multilingual_split_dataset: - n_sequences = truncated_len // (self.context_size) + n_sequences = len(toks) // self.context_size for _ in range(n_sequences): self.runtime_doc_languages.append(self.doc_languages[i]) @@ -493,7 +558,10 @@ def generate_and_save_activations(self, path: str, split_count: int = 10, number # Infer full token count if not provided if number_of_tokens is None: - number_of_tokens = sum(len(example["tokens"]) for example in self.raw_ds) + number_of_tokens = sum( + len(self._truncate_to_context(self._tokens_from_dataset_row(example))) + for example in self.raw_ds + ) number_of_tokens -= number_of_tokens % buffer_size # truncate to full buffers logger.info(f"[ActivationsStore] Using full dataset: {number_of_tokens} tokens") diff --git a/tests/training/test_activation_store_dataset_tokens.py b/tests/training/test_activation_store_dataset_tokens.py new file mode 100644 index 0000000..27e8827 --- /dev/null +++ b/tests/training/test_activation_store_dataset_tokens.py @@ -0,0 +1,220 @@ +from types import SimpleNamespace + +import pytest +import torch + +from clt_forge.config import CLTTrainingRunnerConfig +from clt_forge.training.activations_store import ActivationsStore + + +class DummyTokenizer: + bos_token_id = 0 + + def __call__(self, text: str, add_special_tokens: bool = False): + assert add_special_tokens is False + return {"input_ids": [ord(char) for char in text]} + + +class DummyDataset: + def __init__(self, rows, column_names): + self.rows = rows + self.column_names = column_names + self.renamed = None + + def __len__(self): + return len(self.rows) + + def __getitem__(self, idx): + return self.rows[idx] + + def rename_column(self, old_name: str, new_name: str): + self.renamed = (old_name, new_name) + self.column_names = [ + new_name if column_name == old_name else column_name + for column_name in self.column_names + ] + for row in self.rows: + row[new_name] = row.pop(old_name) + return self + + +def make_store(*, is_dataset_tokenized: bool = True, dataset_text_column: str = "text"): + store = object.__new__(ActivationsStore) + store.cfg = SimpleNamespace( + dataset_path="dummy-dataset", + dataset_text_column=dataset_text_column, + is_dataset_tokenized=is_dataset_tokenized, + is_distributed=False, + is_multilingual_split_dataset=False, + ) + store.model = SimpleNamespace(tokenizer=DummyTokenizer()) + store.context_size = 4 + return store + + +class DummyModel: + def __init__(self, d_in: int = 3, n_layers: int = 1): + self.cfg = SimpleNamespace(n_layers=n_layers) + self.tokenizer = DummyTokenizer() + self.d_in = d_in + self.seen_batches = [] + + def run_with_cache(self, batch_tokens, names_filter, prepend_bos=False): + assert prepend_bos is False + self.seen_batches.append(batch_tokens.detach().cpu()) + base = batch_tokens.float().unsqueeze(-1).repeat(1, 1, self.d_in) + cache = { + name: base + float(i) + for i, name in enumerate(names_filter) + } + return None, cache + + +def make_runner_cfg(*, is_dataset_tokenized: bool, dataset_text_column: str = "text"): + return CLTTrainingRunnerConfig( + device="cpu", + dtype="float32", + model_name="dummy-model", + dataset_path="dummy-dataset", + is_dataset_tokenized=is_dataset_tokenized, + dataset_text_column=dataset_text_column, + d_in=3, + d_latent=4, + train_batch_size_tokens=4, + context_size=4, + n_batches_in_buffer=2, + store_batch_size_prompts=1, + total_training_tokens=8, + n_batches_for_norm_estimate=1, + log_to_wandb=False, + distributed_setup="None", + logger_verbose=False, + ) + + +def test_tokens_from_dataset_row_uses_pre_tokenized_tokens_and_strips_bos(): + store = make_store(is_dataset_tokenized=True) + + toks = store._tokens_from_dataset_row({"tokens": [0, 11, 12, 13]}) + + assert torch.equal(toks, torch.tensor([11, 12, 13])) + + +def test_tokens_from_dataset_row_tokenizes_raw_text(): + store = make_store(is_dataset_tokenized=False) + + toks = store._tokens_from_dataset_row({"text": "abc"}) + + assert torch.equal(toks, torch.tensor([97, 98, 99])) + + +def test_tokens_from_dataset_row_uses_custom_raw_text_column(): + store = make_store(is_dataset_tokenized=False, dataset_text_column="content") + + toks = store._tokens_from_dataset_row({"content": "abc"}) + + assert torch.equal(toks, torch.tensor([97, 98, 99])) + + +def test_validate_dataset_columns_renames_input_ids_for_tokenized_dataset(): + store = make_store(is_dataset_tokenized=True) + store.raw_ds = DummyDataset([{"input_ids": [1, 2, 3, 4]}], ["input_ids"]) + + store._validate_dataset_columns() + + assert store.raw_ds.renamed == ("input_ids", "tokens") + assert store.raw_ds.column_names == ["tokens"] + + +def test_validate_dataset_columns_requires_text_column_for_raw_dataset(): + store = make_store(is_dataset_tokenized=False) + store.raw_ds = DummyDataset([{"content": "abc"}], ["content"]) + + with pytest.raises(ValueError, match="dataset_text_column='text'"): + store._validate_dataset_columns() + + +def test_iterate_raw_dataset_tokens_tokenizes_and_truncates_raw_text(): + store = make_store(is_dataset_tokenized=False) + store.raw_ds = DummyDataset([{"text": "abcdef"}], ["text"]) + + toks = list(store._iterate_raw_dataset_tokens()) + + assert len(toks) == 1 + assert torch.equal(toks[0], torch.tensor([97, 98, 99, 100])) + + +def test_iterate_raw_dataset_tokens_keeps_short_raw_text_sequences(): + store = make_store(is_dataset_tokenized=False) + store.raw_ds = DummyDataset([{"text": "abc"}], ["text"]) + + toks = list(store._iterate_raw_dataset_tokens()) + + assert len(toks) == 1 + assert torch.equal(toks[0], torch.tensor([97, 98, 99])) + + +def test_raw_text_dataset_runs_through_activation_store(monkeypatch): + dataset = DummyDataset([{"text": "abcdefgh"}], ["text"]) + monkeypatch.setattr( + "clt_forge.training.activations_store.load_dataset_auto", + lambda *args, **kwargs: dataset, + ) + cfg = make_runner_cfg(is_dataset_tokenized=False) + model = DummyModel(d_in=cfg.d_in) + + store = ActivationsStore( + model, + cfg, + estimated_norm_scaling_factor_in=torch.ones(model.cfg.n_layers), + estimated_norm_scaling_factor_out=torch.ones(model.cfg.n_layers), + ) + act_in, act_out = next(iter(store)) + + assert act_in.shape == (cfg.train_batch_size_tokens, model.cfg.n_layers, cfg.d_in) + assert act_out.shape == act_in.shape + assert len(model.seen_batches) == cfg.n_batches_in_buffer + assert all(batch.shape == (cfg.store_batch_size_prompts, cfg.context_size + 1) for batch in model.seen_batches) + assert all(torch.all(batch[:, 0] == model.tokenizer.bos_token_id) for batch in model.seen_batches) + + +def test_pre_tokenized_dataset_runs_through_activation_store(monkeypatch): + dataset = DummyDataset([{"tokens": [0, 1, 2, 3, 4, 5, 6, 7, 8]}], ["tokens"]) + monkeypatch.setattr( + "clt_forge.training.activations_store.load_dataset_auto", + lambda *args, **kwargs: dataset, + ) + cfg = make_runner_cfg(is_dataset_tokenized=True) + model = DummyModel(d_in=cfg.d_in) + + store = ActivationsStore( + model, + cfg, + estimated_norm_scaling_factor_in=torch.ones(model.cfg.n_layers), + estimated_norm_scaling_factor_out=torch.ones(model.cfg.n_layers), + ) + act_in, act_out = next(iter(store)) + + assert act_in.shape == (cfg.train_batch_size_tokens, model.cfg.n_layers, cfg.d_in) + assert act_out.shape == act_in.shape + assert len(model.seen_batches) == cfg.n_batches_in_buffer + + +def test_raw_text_dataset_can_generate_cached_activations(monkeypatch, tmp_path): + dataset = DummyDataset([{"text": "abcdefgh"}], ["text"]) + monkeypatch.setattr( + "clt_forge.training.activations_store.load_dataset_auto", + lambda *args, **kwargs: dataset, + ) + cfg = make_runner_cfg(is_dataset_tokenized=False) + model = DummyModel(d_in=cfg.d_in) + store = ActivationsStore( + model, + cfg, + estimated_norm_scaling_factor_in=torch.ones(model.cfg.n_layers), + estimated_norm_scaling_factor_out=torch.ones(model.cfg.n_layers), + ) + + store.generate_and_save_activations(path=str(tmp_path), split_count=1) + + assert (tmp_path / f"ctx_{cfg.context_size}" / "activations_split_0.safetensors").exists()