diff --git a/param_decomp/core/configs.py b/param_decomp/core/configs.py index c82653ca2..5e0980c93 100644 --- a/param_decomp/core/configs.py +++ b/param_decomp/core/configs.py @@ -233,6 +233,7 @@ class SmoothL0ImportanceMinimalityLossConfig(LossMetricConfig): class CIMaskedReconLossConfig(LossMetricConfig, HiddenActsReconstructionMixin): + slow: ClassVar[bool] = False type: Literal["CIMaskedReconLoss"] = "CIMaskedReconLoss" @@ -467,6 +468,15 @@ class ComponentActivationDensityConfig(BaseConfig): ci_alive_threshold: float = 0.0 +class WeightMagnitudeConfig(BaseConfig): + """Per-site `‖V_c‖·‖U_c‖` scatter in descending magnitude order, log y. + + Reads the trained V/U alone — no forward pass and no eval batch.""" + + slow: ClassVar[bool] = True + type: Literal["WeightMagnitude"] = "WeightMagnitude" + + class IdentityCITargetSpec(BaseConfig): """A layer expected to produce an Identity CI pattern over `n_features` features.""" @@ -657,6 +667,7 @@ class UnmaskedNoDeltaReconLossConfig(LossMetricConfig): exists here, and it is non-target-only — the plain and target-pass unions have no member for it.""" + slow: ClassVar[bool] = False type: Literal["UnmaskedNoDeltaReconLoss"] = "UnmaskedNoDeltaReconLoss" diff --git a/param_decomp/core/run.py b/param_decomp/core/run.py index 6fb082f31..66ce802c9 100644 --- a/param_decomp/core/run.py +++ b/param_decomp/core/run.py @@ -333,7 +333,10 @@ def log(self, step: int, record: "LogRecord") -> None: self._last_committed_step = step record = { _METRIC_KEYS.get( - k, f"train/{k}" if k.startswith(("grad_norms/", "loss/", "schedules/")) else k + k, + f"train/{k}" + if k.startswith(("grad_norms/", "loss/", "schedules/", "nontarget_data/")) + else k, ): v for k, v in record.items() } # keys already starting "train/" or "eval/" pass through verbatim diff --git a/param_decomp/core/slow_eval.py b/param_decomp/core/slow_eval.py index 04794f1db..d68499f22 100644 --- a/param_decomp/core/slow_eval.py +++ b/param_decomp/core/slow_eval.py @@ -68,6 +68,7 @@ lower_leaky_hard_sigmoid, upper_leaky_hard_sigmoid, ) +from param_decomp.core.components import ComponentStacks from param_decomp.core.configs import ( DenseCITargetSpec, IdentityCIErrorConfig, @@ -198,6 +199,11 @@ def slow_eval_step( return filter_jit(slow_eval_step, compiler_options=compiler_options) +def _raw_sample(chunks: dict[str, list[np.ndarray]], site: str) -> np.ndarray: + """A site's kept raw values, or empty when `n_batches_accum` kept none.""" + return np.concatenate(chunks[site]) if site in chunks else np.empty(0, np.float32) + + def accumulate_site_reductions( slow_eval_step: SlowEvalStep, model: DecomposedModel, @@ -246,8 +252,8 @@ def accumulate_site_reductions( density_counts=density[site], ci_sums=sums[site], n_positions=total_positions, - lower_sample=np.concatenate(lower_chunks[site]), - preactivations_sample=np.concatenate(preactivations_chunks[site]), + lower_sample=_raw_sample(lower_chunks, site), + preactivations_sample=_raw_sample(preactivations_chunks, site), density_hist=hist.get(site), ) for site in density @@ -578,6 +584,89 @@ def _plot_ci_matrices(matrices: dict[str, np.ndarray], colormap: str, title_pref return _render_figure(fig) +def _component_weight_magnitudes(components: ComponentStacks) -> dict[str, Array]: + return { + name: jnp.linalg.norm(sc.V.astype(jnp.float32), axis=0) + * jnp.linalg.norm(sc.U.astype(jnp.float32), axis=1) + for name, sc in components.sites_items() + } + + +def weight_magnitudes(components: ComponentStacks) -> dict[str, np.ndarray]: + """Per-site `‖V_c‖·‖U_c‖` as host `(C,)` vectors. The norms reduce ON DEVICE, so only + C floats per site cross the boundary — never the V/U matrices themselves.""" + return { + name: np.asarray(value) for name, value in _component_weight_magnitudes(components).items() + } + + +def mean_cis(reductions: dict[str, SiteReduction]) -> dict[str, np.ndarray]: + """Per-site token-weighted mean CI.""" + assert all(r.n_positions > 0 for r in reductions.values()) + return {site: r.ci_sums / r.n_positions for site, r in reductions.items()} + + +def plot_weight_magnitudes(magnitudes: dict[str, np.ndarray]) -> bytes: + """Per-site `‖V_c‖·‖U_c‖` in descending magnitude order, log y. x is a component's rank + within its site, NOT its component id.""" + n_rows, n_cols = _grid_dims(len(magnitudes)) + fig = Figure(figsize=(8 * n_cols, 3 * n_rows)) + axs = fig.subplots(n_rows, n_cols, squeeze=False) + flat_axes = axs.T.ravel() + for ax in flat_axes[len(magnitudes) :]: + ax.set_visible(False) + for ax, (name, values) in zip(flat_axes, magnitudes.items(), strict=False): + ax.scatter(range(len(values)), np.sort(values)[::-1], marker="x", s=10) + ax.set_yscale("log") + ax.set_xlabel("Component (descending ‖V‖·‖U‖)") + ax.set_ylabel("‖V‖·‖U‖") + ax.set_title(name, fontsize=10) + fig.tight_layout() + return _render_figure(fig) + + +def plot_mean_component_cis_two_streams( + target_mean_cis: dict[str, np.ndarray], + nontarget_mean_cis: dict[str, np.ndarray], +) -> tuple[bytes, bytes]: + """Both streams' mean CI on one axis per site, ordered by descending TARGET mean. + + The nontarget series takes the same permutation rather than its own, so a component's + two series line up vertically.""" + assert target_mean_cis.keys() == nontarget_mean_cis.keys(), ( + sorted(target_mean_cis), + sorted(nontarget_mean_cis), + ) + n_rows, n_cols = _grid_dims(len(target_mean_cis)) + ordered = { + name: (target[order], nontarget_mean_cis[name][order]) + for name, target in target_mean_cis.items() + for order in [np.argsort(target)[::-1]] + } + images: list[bytes] = [] + for log_y in (False, True): + fig = Figure(figsize=(8 * n_cols, 3 * n_rows)) + axs = fig.subplots(n_rows, n_cols, squeeze=False) + flat_axes = axs.T.ravel() + for ax in flat_axes[len(ordered) :]: + ax.set_visible(False) + for ax, (name, (target, nontarget)) in zip(flat_axes, ordered.items(), strict=False): + x = np.arange(len(target)) + if log_y: + ax.set_yscale("log") + ax.fill_between(x, target, step="mid", color="#1f77b4", label="target") + ax.fill_between( + x, nontarget, step="mid", color="#d62728", label="non-target", alpha=0.6 + ) + ax.set_xlabel("Component (sorted by target mean CI)") + ax.set_ylabel("mean CI") + ax.set_title(name, fontsize=10) + ax.legend(fontsize=7) + fig.tight_layout() + images.append(_render_figure(fig)) + return images[0], images[1] + + def plot_permuted_ci_heatmaps( position_ci: dict[str, PositionCI], permutation: dict[str, "Literal['identity', 'dense']"] ) -> tuple[bytes, bytes]: diff --git a/param_decomp/core/tests/test_eval_averaging_parity.py b/param_decomp/core/tests/test_eval_averaging_parity.py index ef898f62b..9bf2944bb 100644 --- a/param_decomp/core/tests/test_eval_averaging_parity.py +++ b/param_decomp/core/tests/test_eval_averaging_parity.py @@ -132,12 +132,14 @@ def step( jnp.array([0, 0], dtype=jnp.uint32), train_steps=0, eval_steps=2, + stream="nontarget", ) context = LMEvalContext( state=_state_stub(), # pyright: ignore[reportArgumentType] now_step=0, pass_index=0, batches=(jnp.asarray(1.0), jnp.asarray(3.0)), + target_batches=None, ) assert operation.run(context)["eval/loss/probe/hidden_acts_reconstruction"] == 2.0 diff --git a/param_decomp/core/train.py b/param_decomp/core/train.py index 99348b1ad..c3fd3c9f0 100644 --- a/param_decomp/core/train.py +++ b/param_decomp/core/train.py @@ -1185,14 +1185,18 @@ def make_targeted_train_step[PreparedT]( for term in objective.target.recon if term.hidden_acts_reconstruction is not None }, - "nontarget/impmin": objective.nontarget.impmin_coeff, - **{f"nontarget/{term.name}": term.coeff for term in nt_terms}, } if objective.target.imp.cfg.frequency is not None: coeff_schedules[f"{objective.target.imp.name}/frequency"] = ( objective.target.imp.cfg.frequency.coeff ) + imp_name = objective.target.imp.name + nontarget_coeff_schedules: dict[str, LossCoeff] = { + imp_name: objective.nontarget.impmin_coeff, + **{term.name: term.coeff for term in nt_terms}, + } + def nontarget_draw_loss( model: DecomposedModel[PreparedT], prepared_weights: PreparedT, @@ -1360,8 +1364,8 @@ def loss_fn( nt_imp_lp, nt_imp_freq = imp_min_terms(nt_ci.upper, atoms.imp_min, imp_min_param) nt_total = nt_imp_coeff * nt_imp_lp + freq_coeff * nt_imp_freq nt_aux = { - f"loss/nontarget/{atoms.imp_loss_key}": nt_imp_lp, - "loss/nontarget/freq": nt_imp_freq, + f"nontarget_data/loss/{imp_name}": nt_imp_lp, + "nontarget_data/loss/FrequencyMinimalityLoss": nt_imp_freq, } nt_breakdowns = atoms.grid_losses( nt_terms, @@ -1374,8 +1378,8 @@ def loss_fn( nt_terms, nt_recon_coeffs, nt_breakdowns, strict=True ): nt_total = nt_total + coeff * breakdown.total - nt_aux[f"loss/nontarget/{term.name}"] = breakdown.total - nt_aux["loss/nontarget/total"] = nt_total + nt_aux[f"nontarget_data/loss/{term.name}"] = breakdown.total + nt_aux["nontarget_data/loss/total"] = nt_total total_loss = total_loss + nt_total reported_total = reported_total + nt_total return total_loss, (reported_total, imp_lp, imp_freq, term_breakdowns, nt_aux) @@ -1446,6 +1450,12 @@ def loss_fn( | nt_aux | wd_metrics | _scheduled_coeff_metrics(step_f32, atoms.total_steps, coeff_schedules) + | { + f"nontarget_data/{key}": value + for key, value in _scheduled_coeff_metrics( + step_f32, atoms.total_steps, nontarget_coeff_schedules + ).items() + } ) return new_state, metrics diff --git a/param_decomp/experiments/eval_config.py b/param_decomp/experiments/eval_config.py index 4545dc296..09add312c 100644 --- a/param_decomp/experiments/eval_config.py +++ b/param_decomp/experiments/eval_config.py @@ -14,6 +14,7 @@ CI_L0Config, CIHiddenActsReconLossConfig, CIHistogramsConfig, + CIMaskedReconLossConfig, CIMeanPerComponentConfig, ComponentActivationDensityConfig, HiddenActsReconstructionMixin, @@ -22,7 +23,9 @@ PermutedCIPlotsConfig, PGDReconLossConfig, StochasticHiddenActsReconLossConfig, + UnmaskedNoDeltaReconLossConfig, UVPlotsConfig, + WeightMagnitudeConfig, WellTemperednessConfig, ) from param_decomp.core.eval_schedule import EvalSchedule, Every, FirstThenEvery @@ -31,6 +34,7 @@ CEandKLLossesConfig, CIMaskedAttnPatternsReconLossConfig, StochasticAttnPatternsReconLossConfig, + TwoStreamCIMeanPerComponentConfig, ) AnyEvalMetricConfig = Annotated[ @@ -40,6 +44,7 @@ | CIHistogramsConfig | CI_L0Config | CIMaskedAttnPatternsReconLossConfig + | CIMaskedReconLossConfig | CIMeanPerComponentConfig | ComponentActivationDensityConfig | IdentityCIErrorConfig @@ -47,7 +52,10 @@ | PGDReconLossConfig | StochasticAttnPatternsReconLossConfig | StochasticHiddenActsReconLossConfig + | TwoStreamCIMeanPerComponentConfig + | UnmaskedNoDeltaReconLossConfig | UVPlotsConfig + | WeightMagnitudeConfig | WellTemperednessConfig, Discriminator("type"), ] diff --git a/param_decomp/experiments/lm/diagnostic_eval_operations.py b/param_decomp/experiments/lm/diagnostic_eval_operations.py index da648c1ec..b4e794a75 100644 --- a/param_decomp/experiments/lm/diagnostic_eval_operations.py +++ b/param_decomp/experiments/lm/diagnostic_eval_operations.py @@ -6,6 +6,7 @@ import numpy as np from jaxtyping import PRNGKeyArray +from param_decomp.core.ci_fn import CIFn from param_decomp.core.configs import ( CIHiddenActsReconLossConfig, CIHistogramsConfig, @@ -40,9 +41,13 @@ compute_identity_ci_errors, make_position_ci_step, make_slow_eval_step, + mean_cis, + plot_mean_component_cis_two_streams, + plot_weight_magnitudes, render_permutation_figures, render_slow_eval_figures, resolve_permutation_metrics, + weight_magnitudes, ) from param_decomp.experiments.lm.attn_patterns_eval import ( accumulate_attn_patterns, @@ -56,6 +61,11 @@ ) from param_decomp.experiments.lm.eval_context import LMEvalContext from param_decomp.experiments.lm.eval_keys import EvalKeyStream +from param_decomp.experiments.lm.scalar_eval_operations import ( + Stream, + stream_batches, + stream_log_prefix, +) def _render_selected_figures( @@ -94,6 +104,7 @@ def make_attention_operation( run_key: PRNGKeyArray, train_steps: int, compiler_options: dict[str, bool | int | str], + stream: Stream, ) -> EvalOperation[LMEvalContext]: match metric: case CIMaskedAttnPatternsReconLossConfig(): @@ -109,13 +120,14 @@ def run(context: LMEvalContext) -> LogRecord: model, context.state.decomposition.components, context.state.decomposition.ci_fn, - list(context.batches), + list(stream_batches(stream, context)), jax.random.fold_in( run_key, EvalKeyStream.ATTENTION_PATTERNS * train_steps + context.pass_index ), ) + prefix = stream_log_prefix(stream, context) return { - f"eval/loss/{name}": value + f"{prefix}loss/{name}": value for name, value in attn_patterns_log_entries(metric.type, reductions).items() } @@ -130,6 +142,7 @@ def make_hidden_acts_operation( run_key: PRNGKeyArray, train_steps: int, compiler_options: dict[str, bool | int | str], + stream: Stream, ) -> EvalOperation[LMEvalContext]: match metric: case CIHiddenActsReconLossConfig(): @@ -145,19 +158,89 @@ def run(context: LMEvalContext) -> LogRecord: model, context.state.decomposition.components, context.state.decomposition.ci_fn, - list(context.batches), + list(stream_batches(stream, context)), jax.random.fold_in( run_key, EvalKeyStream.HIDDEN_ACTS * train_steps + context.pass_index ), ) + prefix = stream_log_prefix(stream, context) return { - f"eval/slow/loss/{name}": value + f"{prefix}slow/loss/{name}": value for name, value in hidden_acts_log_entries(metric.type, reductions).items() } return EvalOperation(schedule, run) +def _render_weight_magnitudes( + magnitudes: dict[str, np.ndarray], now_step: int +) -> DeferredMediaRecord: + return DeferredMediaRecord( + step_key="slow_eval/figure_step", + step=now_step, + media={"slow_eval/figures/weight_magnitude": plot_weight_magnitudes(magnitudes)}, + ) + + +def make_weight_magnitude_operation( + schedule: EvalSchedule, renderer: BackgroundRenderer +) -> EvalOperation[LMEvalContext]: + """`‖V_c‖·‖U_c‖` per site. Reads the trained V/U only — no model, no batch, no step.""" + + def run(context: LMEvalContext) -> LogRecord: + magnitudes = weight_magnitudes(context.state.decomposition.components) + renderer.submit(partial(_render_weight_magnitudes, magnitudes, context.now_step)) + return {} + + return EvalOperation(schedule, run) + + +def _render_two_stream_ci_means( + target: dict[str, np.ndarray], + nontarget: dict[str, np.ndarray], + now_step: int, +) -> DeferredMediaRecord: + linear, log = plot_mean_component_cis_two_streams(target, nontarget) + return DeferredMediaRecord( + step_key="slow_eval/figure_step", + step=now_step, + media={ + "slow_eval/figures/ci_mean_per_component_two_streams": linear, + "slow_eval/figures/ci_mean_per_component_two_streams_log": log, + }, + ) + + +def make_two_stream_ci_mean_operation( + schedule: EvalSchedule, + model: DecomposedModel, + ci_capture_keys: CaptureKeys, + compiler_options: dict[str, bool | int | str], + renderer: BackgroundRenderer, +) -> EvalOperation[LMEvalContext]: + """Both streams' mean CI per component in one figure, ordered by the target mean.""" + step = make_slow_eval_step(model, ci_capture_keys, 0.0, None, compiler_options) + + def stream_mean_cis(ci_fn: CIFn, batches: tuple[jax.Array, ...]) -> dict[str, np.ndarray]: + return mean_cis( + accumulate_site_reductions(step, model, ci_fn, list(batches), n_batches_accum=0) + ) + + def run(context: LMEvalContext) -> LogRecord: + ci_fn = context.state.decomposition.ci_fn + renderer.submit( + partial( + _render_two_stream_ci_means, + stream_mean_cis(ci_fn, stream_batches("target", context)), + stream_mean_cis(ci_fn, stream_batches("nontarget", context)), + context.now_step, + ) + ) + return {} + + return EvalOperation(schedule, run) + + def make_site_figures_operation( metric: CIHistogramsConfig | ComponentActivationDensityConfig | CIMeanPerComponentConfig, schedule: EvalSchedule, @@ -165,6 +248,7 @@ def make_site_figures_operation( ci_capture_keys: CaptureKeys, compiler_options: dict[str, bool | int | str], renderer: BackgroundRenderer, + stream: Stream, ) -> EvalOperation[LMEvalContext]: match metric: case CIHistogramsConfig(): @@ -196,7 +280,7 @@ def run(context: LMEvalContext) -> LogRecord: step, model, context.state.decomposition.ci_fn, - list(context.batches), + list(stream_batches(stream, context)), limit, ) renderer.submit(partial(_render_selected_figures, reductions, wanted, context.now_step)) @@ -212,6 +296,7 @@ def make_permutation_operation( ci_capture_keys: CaptureKeys, compiler_options: dict[str, bool | int | str], renderer: BackgroundRenderer, + stream: Stream, ) -> EvalOperation[LMEvalContext]: spec = resolve_permutation_metrics(model.site_names, [metric]) position_step = make_position_ci_step(model, ci_capture_keys, compiler_options) @@ -221,12 +306,13 @@ def run(context: LMEvalContext) -> LogRecord: position_step, model, context.state.decomposition.ci_fn, - list(context.batches), + list(stream_batches(stream, context)), ) match metric: case IdentityCIErrorConfig(): errors = compute_identity_ci_errors(spec, position_ci, IDENTITY_CI_ERROR_TOLERANCE) - return {f"eval/slow/{name}": value for name, value in errors.items()} + prefix = stream_log_prefix(stream, context) + return {f"{prefix}slow/{name}": value for name, value in errors.items()} case UVPlotsConfig(): include_ci_heatmaps = False components = { diff --git a/param_decomp/experiments/lm/eval.py b/param_decomp/experiments/lm/eval.py index fdd09b804..0ea42786c 100644 --- a/param_decomp/experiments/lm/eval.py +++ b/param_decomp/experiments/lm/eval.py @@ -42,6 +42,7 @@ import math from collections.abc import Callable, Mapping from dataclasses import dataclass +from typing import Literal import jax import jax.numpy as jnp @@ -204,6 +205,49 @@ def _ce[PreparedT](batch: _PreparedLMBatch[PreparedT], logits: Array) -> Array: return _row_masked_cross_entropy(logits, batch.tokens, batch.valid_row_mask) +type MaskingArm = Literal["ci_masked", "unmasked"] +"""A masking arm authorable as an eval on its own. Both pin every weight-delta mask to +zero.""" + + +def make_masked_kl_step[PreparedT]( + model_static: DecomposedModel[PreparedT], + ci_capture_keys: CaptureKeys, + arm: MaskingArm, + mesh: Mesh | None = None, + compiler_options: dict[str, bool | int | str] | None = None, + *, + n_valid_rows: int | None = None, +) -> ScalarStep: + """KL against the target output under ONE masking arm — one clean forward, one masked. + + The key is the spelling `CEandKLLosses` reports this arm under, so the two are one + quantity under one name.""" + assert model_static.has_position_axis, "masked KL is LM-only and requires a position axis" + + def eval_step( + model: DecomposedModel[PreparedT], + components: ComponentStacks, + ci_fn: CIFn, + token_ids: Array, + key: PRNGKeyArray, + ) -> dict[str, Array]: + del key # neither arm draws masks + batch = _prepare_lm_batch( + model, components, ci_fn, token_ids, mesh, n_valid_rows, ci_capture_keys + ) + match arm: + case "ci_masked": + masks = batch.ci_lower + case "unmasked": + masks = {site: jnp.ones_like(batch.ci_lower[site]) for site in model.site_names} + zeros_delta = {site: jnp.zeros(batch.tokens.shape, COMPUTE_DT) for site in model.site_names} + logits = _compute_masked_output(model, batch, masks, zeros_delta, mesh, frozenset()) + return {f"ce_kl/kl_{arm}": _kl(batch, logits)} + + return filter_jit(eval_step, compiler_options=compiler_options) + + def make_ce_kl_step[PreparedT]( model_static: DecomposedModel[PreparedT], ci_capture_keys: CaptureKeys, diff --git a/param_decomp/experiments/lm/eval_config.py b/param_decomp/experiments/lm/eval_config.py index b850743d7..09722b61a 100644 --- a/param_decomp/experiments/lm/eval_config.py +++ b/param_decomp/experiments/lm/eval_config.py @@ -62,3 +62,14 @@ class ArithmeticCIGridConfig(BaseConfig): b_range: tuple[int, int] = (1, 100) thresholds: list[float] = Field(default_factory=lambda: [0.1]) top_k: PositiveInt = 24 + + +class TwoStreamCIMeanPerComponentConfig(BaseConfig): + """Both streams' mean CI per component on ONE axis per site, ordered by descending + TARGET mean and coloured by stream. + + Computes `CIMeanPerComponent`'s reduction on both streams, so authoring both pays for + the nontarget pass twice. Refuses on a plain run, which has no target stream.""" + + slow: ClassVar[bool] = True + type: Literal["TwoStreamCIMeanPerComponent"] = "TwoStreamCIMeanPerComponent" diff --git a/param_decomp/experiments/lm/eval_context.py b/param_decomp/experiments/lm/eval_context.py index cc674bafa..e521eb79f 100644 --- a/param_decomp/experiments/lm/eval_context.py +++ b/param_decomp/experiments/lm/eval_context.py @@ -11,3 +11,7 @@ class LMEvalContext(EvalInvocation): pass_index: int batches: tuple[jax.Array, ...] + target_batches: tuple[jax.Array, ...] | None = None + """A tPD run's target-stream draws; `None` on a plain run, which has no second stream. + That `None` is also what tells every log key which run kind it is in + (`scalar_eval_operations.stream_log_prefix`).""" diff --git a/param_decomp/experiments/lm/eval_operations.py b/param_decomp/experiments/lm/eval_operations.py index 0917f4625..1c8111b3c 100644 --- a/param_decomp/experiments/lm/eval_operations.py +++ b/param_decomp/experiments/lm/eval_operations.py @@ -1,5 +1,7 @@ """LM evaluation operation binding and execution.""" +from collections.abc import Callable + import jax import numpy as np from jax.sharding import Mesh, NamedSharding @@ -10,15 +12,19 @@ CI_L0Config, CIHiddenActsReconLossConfig, CIHistogramsConfig, + CIMaskedReconLossConfig, CIMeanPerComponentConfig, ComponentActivationDensityConfig, IdentityCIErrorConfig, PermutedCIPlotsConfig, PGDReconLossConfig, StochasticHiddenActsReconLossConfig, + UnmaskedNoDeltaReconLossConfig, UVPlotsConfig, + WeightMagnitudeConfig, WellTemperednessConfig, ) +from param_decomp.core.eval_schedule import EvalSchedule from param_decomp.core.model import DecomposedModel from param_decomp.core.run import ( BackgroundRenderer, @@ -35,20 +41,26 @@ make_hidden_acts_operation, make_permutation_operation, make_site_figures_operation, + make_two_stream_ci_mean_operation, + make_weight_magnitude_operation, ) from param_decomp.experiments.lm.eval_config import ( ArithmeticCIGridConfig, CEandKLLossesConfig, CIMaskedAttnPatternsReconLossConfig, StochasticAttnPatternsReconLossConfig, + TwoStreamCIMeanPerComponentConfig, ) from param_decomp.experiments.lm.eval_context import LMEvalContext from param_decomp.experiments.lm.eval_keys import EvalKeyStream from param_decomp.experiments.lm.resolved import LMAnyRun from param_decomp.experiments.lm.scalar_eval_operations import ( + Stream, make_ce_kl_operation, make_ci_l0_operation, make_fresh_pgd_operation, + make_masked_kl_operation, + stream_batches, ) from param_decomp.infra.dataset_store import read_dataset_meta from param_decomp.pretrain.batch_data import BatchSchedule, ShardServer, scan_shards @@ -68,8 +80,11 @@ def make_lm_evaluation( n_proc: int, sink: MetricsSink, compiler_options: dict[str, bool | int | str], + target_pool_batches_for: Callable[[int], list[jax.Array]] | None = None, ) -> Evaluation[LMEvalContext]: - """Construct one executable operation for every authored LM metric.""" + """Construct one executable operation for every authored LM metric (for tPD, one per + stream the metric measures; `target_pool_batches_for` draws the target stream, and + `None` marks a plain run).""" pd = built.pd capture_inputs = built.ci_fn.capture_keys data = built.data @@ -87,102 +102,158 @@ def batches(pass_index: int) -> list[jax.Array]: for j in range(eval.n_steps) ] + targeted = target_pool_batches_for is not None + all_streams: tuple[Stream, ...] = ("nontarget", "target") if targeted else ("nontarget",) + primary_stream: Stream = "target" if targeted else "nontarget" + def well_temperedness_inputs( context: LMEvalContext, ) -> tuple[jax.Array, PRNGKeyArray]: - return context.batches[0], jax.random.fold_in( + return stream_batches(primary_stream, context)[0], jax.random.fold_in( run_key, EvalKeyStream.WELL_TEMPEREDNESS * pd.steps + context.pass_index ) - def make_operation(metric: AnyEvalMetricConfig) -> EvalOperation[LMEvalContext]: + def per_stream( + maker: Callable[..., EvalOperation[LMEvalContext]], + metric: object, + schedule: EvalSchedule, + streams: tuple[Stream, ...], + ) -> tuple[EvalOperation[LMEvalContext], ...]: + return tuple( + maker( + metric, + schedule, + stream, + model, + capture_inputs, + run_key, + pd.steps, + eval.n_steps, + mesh, + compiler_options, + ) + for stream in streams + ) + + def make_operations(metric: AnyEvalMetricConfig) -> tuple[EvalOperation[LMEvalContext], ...]: schedule = schedule_for(metric, eval) match metric: case CEandKLLossesConfig(): - return make_ce_kl_operation( - metric, - schedule, - model, - capture_inputs, - run_key, - pd.steps, - eval.n_steps, - mesh, - compiler_options, - ) + return per_stream(make_ce_kl_operation, metric, schedule, all_streams) + case CIMaskedReconLossConfig(): + return per_stream(make_masked_kl_operation, "ci_masked", schedule, all_streams) + case UnmaskedNoDeltaReconLossConfig(): + return per_stream(make_masked_kl_operation, "unmasked", schedule, (primary_stream,)) case CI_L0Config(): - return make_ci_l0_operation( - metric, - schedule, - model, - capture_inputs, - run_key, - pd.steps, - eval.n_steps, - mesh, - compiler_options, - ) + return per_stream(make_ci_l0_operation, metric, schedule, all_streams) case PGDReconLossConfig(): - return make_fresh_pgd_operation( - metric, - schedule, - model, - capture_inputs, - run_key, - pd.steps, - eval.n_steps, - mesh, - compiler_options, - ) + return per_stream(make_fresh_pgd_operation, metric, schedule, all_streams) case CIMaskedAttnPatternsReconLossConfig() | StochasticAttnPatternsReconLossConfig(): - return make_attention_operation( - metric, schedule, model, capture_inputs, run_key, pd.steps, compiler_options + return ( + make_attention_operation( + metric, + schedule, + model, + capture_inputs, + run_key, + pd.steps, + compiler_options, + primary_stream, + ), ) case CIHiddenActsReconLossConfig() | StochasticHiddenActsReconLossConfig(): - return make_hidden_acts_operation( - metric, schedule, model, capture_inputs, run_key, pd.steps, compiler_options + return ( + make_hidden_acts_operation( + metric, + schedule, + model, + capture_inputs, + run_key, + pd.steps, + compiler_options, + primary_stream, + ), ) case ( CIHistogramsConfig() | ComponentActivationDensityConfig() | CIMeanPerComponentConfig() ): - return make_site_figures_operation( - metric, schedule, model, capture_inputs, compiler_options, renderer + return ( + make_site_figures_operation( + metric, + schedule, + model, + capture_inputs, + compiler_options, + renderer, + primary_stream, + ), ) case PermutedCIPlotsConfig() | UVPlotsConfig() | IdentityCIErrorConfig(): - return make_permutation_operation( - metric, schedule, model, capture_inputs, compiler_options, renderer + return ( + make_permutation_operation( + metric, + schedule, + model, + capture_inputs, + compiler_options, + renderer, + primary_stream, + ), ) case WellTemperednessConfig(): - return make_well_temperedness_operation( - metric, - schedule, - model, - capture_inputs, - mesh, - compiler_options, - inputs_for_context=well_temperedness_inputs, - figure_rendering=renderer if sink.accepts_deferred_media else None, + return ( + make_well_temperedness_operation( + metric, + schedule, + model, + capture_inputs, + mesh, + compiler_options, + inputs_for_context=well_temperedness_inputs, + figure_rendering=renderer if sink.accepts_deferred_media else None, + ), ) case ArithmeticCIGridConfig(): - return make_arithmetic_operation( - metric, - schedule, - built.target, - model, - capture_inputs, - mesh, - n_proc, - sink, - run_key, - pd.steps, - compiler_options, + return ( + make_arithmetic_operation( + metric, + schedule, + built.target, + model, + capture_inputs, + mesh, + n_proc, + sink, + run_key, + pd.steps, + compiler_options, + ), + ) + case TwoStreamCIMeanPerComponentConfig(): + return ( + make_two_stream_ci_mean_operation( + schedule, model, capture_inputs, compiler_options, renderer + ), ) + case WeightMagnitudeConfig(): + return (make_weight_magnitude_operation(schedule, renderer),) - operations = tuple(make_operation(metric) for metric in eval.metrics) + authored = {type(metric) for metric in eval.metrics} + assert targeted or TwoStreamCIMeanPerComponentConfig not in authored, ( + "TwoStreamCIMeanPerComponent measures the target stream; a plain run has none" + ) + assert not ( + authored & {CIMaskedReconLossConfig, UnmaskedNoDeltaReconLossConfig} + and CEandKLLossesConfig in authored + ), "the single-arm KL evals emit keys CEandKLLosses also emits; author one or the other" + operations = tuple( + operation for metric in eval.metrics for operation in make_operations(metric) + ) def make_context(state: TrainState, now_step: int) -> LMEvalContext: pass_index = now_step // eval.every @@ -191,6 +262,11 @@ def make_context(state: TrainState, now_step: int) -> LMEvalContext: now_step=now_step, pass_index=pass_index, batches=tuple(batches(pass_index)), + target_batches=( + None + if target_pool_batches_for is None + else tuple(target_pool_batches_for(pass_index)) + ), ) return Evaluation(operations, make_context) diff --git a/param_decomp/experiments/lm/scalar_eval_operations.py b/param_decomp/experiments/lm/scalar_eval_operations.py index 6a7425728..9cb3ae320 100644 --- a/param_decomp/experiments/lm/scalar_eval_operations.py +++ b/param_decomp/experiments/lm/scalar_eval_operations.py @@ -1,5 +1,7 @@ """Independent CE/KL, causal-L0, and fresh-PGD LM operations.""" +from typing import Literal + import jax.numpy as jnp from jax import random from jax.sharding import Mesh @@ -13,15 +15,41 @@ from param_decomp.core.recon_eval import FreshPGDReconEval from param_decomp.core.run import EvalOperation from param_decomp.experiments.lm.eval import ( + MaskingArm, ScalarStep, make_ce_kl_step, make_ci_l0_step, make_fresh_pgd_step, + make_masked_kl_step, ) from param_decomp.experiments.lm.eval_config import CEandKLLossesConfig from param_decomp.experiments.lm.eval_context import LMEvalContext from param_decomp.experiments.lm.eval_keys import EvalKeyStream +type Stream = Literal["nontarget", "target"] +"""Which STREAM an eval operation measures. ONE value, not a (batch source, log prefix) +pair, so target-stream batches cannot be spelled under the nontarget stream's log keys.""" + + +def stream_batches(stream: Stream, context: LMEvalContext) -> tuple[Array, ...]: + match stream: + case "nontarget": + return context.batches + case "target": + assert context.target_batches is not None, ( + "target-stream metrics need a tPD run's prompt pool; a plain run has none" + ) + return context.target_batches + + +def stream_log_prefix(stream: Stream, context: LMEvalContext) -> str: + targeted = context.target_batches is not None + match stream: + case "nontarget": + return "eval/nontarget_data/" if targeted else "eval/" + case "target": + return "eval/" + def _make_scalar_operation( schedule: EvalSchedule, @@ -31,10 +59,12 @@ def _make_scalar_operation( run_key: PRNGKeyArray, train_steps: int, eval_steps: int, + stream: Stream, ) -> EvalOperation[LMEvalContext]: def run(context: LMEvalContext) -> LogRecord: + log_prefix = stream_log_prefix(stream, context) sums: dict[str, Array] = {} - for batch_index, tokens in enumerate(context.batches): + for batch_index, tokens in enumerate(stream_batches(stream, context)): key = random.fold_in( run_key, EvalKeyStream.SCALARS * train_steps + context.pass_index * eval_steps + batch_index, @@ -49,7 +79,7 @@ def run(context: LMEvalContext) -> LogRecord: for name, value in values.items(): if name.startswith(prefixes): sums[name] = sums.get(name, jnp.zeros(())) + value - return {f"eval/{name}": float(value) / eval_steps for name, value in sums.items()} + return {f"{log_prefix}{name}": float(value) / eval_steps for name, value in sums.items()} return EvalOperation(schedule, run) @@ -57,6 +87,7 @@ def run(context: LMEvalContext) -> LogRecord: def make_ce_kl_operation( metric: CEandKLLossesConfig, schedule: EvalSchedule, + stream: Stream, model: DecomposedModel, ci_capture_keys: CaptureKeys, run_key: PRNGKeyArray, @@ -65,7 +96,7 @@ def make_ce_kl_operation( mesh: Mesh, compiler_options: dict[str, bool | int | str], ) -> EvalOperation[LMEvalContext]: - scalars = _make_scalar_operation( + return _make_scalar_operation( schedule, make_ce_kl_step(model, ci_capture_keys, metric.rounding_threshold, mesh, compiler_options), ("ce_kl/",), @@ -73,14 +104,40 @@ def make_ce_kl_operation( run_key, train_steps, eval_steps, + stream, ) - return scalars + +def make_masked_kl_operation( + arm: MaskingArm, + schedule: EvalSchedule, + stream: Stream, + model: DecomposedModel, + ci_capture_keys: CaptureKeys, + run_key: PRNGKeyArray, + train_steps: int, + eval_steps: int, + mesh: Mesh, + compiler_options: dict[str, bool | int | str], +) -> EvalOperation[LMEvalContext]: + """ONE masking arm, authored as the loss config that names the same construction — + `CIMaskedReconLoss` / `UnmaskedNoDeltaReconLoss`, as `PGDReconLoss` already is.""" + return _make_scalar_operation( + schedule, + make_masked_kl_step(model, ci_capture_keys, arm, mesh, compiler_options), + (f"ce_kl/kl_{arm}",), + model, + run_key, + train_steps, + eval_steps, + stream, + ) def make_ci_l0_operation( metric: CI_L0Config, schedule: EvalSchedule, + stream: Stream, model: DecomposedModel, ci_capture_keys: CaptureKeys, run_key: PRNGKeyArray, @@ -104,12 +161,14 @@ def make_ci_l0_operation( run_key, train_steps, eval_steps, + stream, ) def run(context: LMEvalContext) -> LogRecord: record = dict(scalars.run(context)) - prefix = f"eval/l0/{metric.ci_alive_threshold}_" - record["eval/l0/bar_chart"] = BarChart( + log_prefix = stream_log_prefix(stream, context) + prefix = f"{log_prefix}l0/{metric.ci_alive_threshold}_" + record[f"{log_prefix}l0/bar_chart"] = BarChart( rows=tuple( (name.removeprefix(prefix), value) for name, value in record.items() @@ -127,6 +186,7 @@ def run(context: LMEvalContext) -> LogRecord: def make_fresh_pgd_operation( metric: PGDReconLossConfig, schedule: EvalSchedule, + stream: Stream, model: DecomposedModel, ci_capture_keys: CaptureKeys, run_key: PRNGKeyArray, @@ -150,4 +210,5 @@ def make_fresh_pgd_operation( run_key, train_steps, eval_steps, + stream, ) diff --git a/param_decomp/experiments/lm/training_targeted.py b/param_decomp/experiments/lm/training_targeted.py index 2019c4288..8175f85bb 100644 --- a/param_decomp/experiments/lm/training_targeted.py +++ b/param_decomp/experiments/lm/training_targeted.py @@ -119,12 +119,15 @@ def train_targeted( server.per_process, jax.local_device_count(), ) - per_process_target = target_batch // n_proc + + def pool_global_batch(seed: int, step: int, batch: int) -> jax.Array: + per_process = batch // n_proc + rows = pool_batch(pool, seed, step, batch) + local = rows[jax.process_index() * per_process :][:per_process] + return global_token_batch(local, mesh, batch) def sample_target_batch(step: int) -> jax.Array: - rows = pool_batch(pool, built.pd.seed, step, target_batch) - local = rows[jax.process_index() * per_process_target :][:per_process_target] - return global_token_batch(local, mesh, target_batch) + return pool_global_batch(built.pd.seed, step, target_batch) def sample_nontarget_batch(step: int) -> jax.Array: return global_token_batch(server.local_batch(step), mesh, nontarget_batch) @@ -136,8 +139,27 @@ def sample_nontarget_batch(step: int) -> jax.Array: "eval must land on a train-log step: the tok/s window resets after eval, so a " "mid-window eval would corrupt the next step-time estimate" ) + eval_target_batch = eval_config.batch_size + + def eval_target_pool_batches(pass_index: int) -> list[jax.Array]: + """The eval pass's target stream: training's pure `(seed, step)` pool sampler on + the `seed + 1` stream, so an eval never scores the rows the step just trained.""" + n_batches = eval_config.n_steps + return [ + pool_global_batch(built.pd.seed + 1, pass_index * n_batches + j, eval_target_batch) + for j in range(n_batches) + ] + evaluation = make_lm_evaluation( - built, eval_config, model, run_key, mesh, n_proc, sink, cfg.runtime.compiler_options + built, + eval_config, + model, + run_key, + mesh, + n_proc, + sink, + cfg.runtime.compiler_options, + target_pool_batches_for=eval_target_pool_batches, ) run_targeted_decomposition_training( diff --git a/param_decomp/experiments/tms/test_targeted_tms.py b/param_decomp/experiments/tms/test_targeted_tms.py index c1073bef0..a5afc5d5e 100644 --- a/param_decomp/experiments/tms/test_targeted_tms.py +++ b/param_decomp/experiments/tms/test_targeted_tms.py @@ -130,8 +130,8 @@ def test_targeted_two_pass_step_trains(): assert "faith" not in metrics # Both passes' losses are reported. assert "loss/StochasticReconLoss" in metrics - assert "loss/nontarget/StochasticReconLoss" in metrics - assert "loss/nontarget/total" in metrics + assert "nontarget_data/loss/StochasticReconLoss" in metrics + assert "nontarget_data/loss/total" in metrics def test_targeted_step_trains_with_persistent_adversary(): @@ -245,7 +245,7 @@ def all_ones_recon_at_delta(delta_value: float) -> jax.Array: delta_on = all_ones_recon_at_delta(1.0) _, metrics = step(model, state, target_batch, nontarget_batch, jax.random.PRNGKey(300)) - reported = metrics["loss/nontarget/UnmaskedNoDeltaReconLoss"] + reported = metrics["nontarget_data/loss/UnmaskedNoDeltaReconLoss"] assert jnp.allclose(reported, delta_off, rtol=1e-4, atol=1e-7) # Delta ON would make the same all-ones forward reproduce the frozen output, so its # loss collapses; a material gap pins that the reported arm really ran delta-off. @@ -635,5 +635,5 @@ def sample_nontarget_batch(step: int) -> jax.Array: lines = (run_dir / "metrics.jsonl").read_text().strip().splitlines() assert lines, "the targeted run logged no metrics" last = json.loads(lines[-1]) - assert "train/loss/nontarget/total" in last + assert "train/nontarget_data/loss/total" in last assert not any("Faithfulness" in k for k in last) diff --git a/param_decomp/tests/test_eval_tier.py b/param_decomp/tests/test_eval_tier.py index 132fb80e8..965def1c8 100644 --- a/param_decomp/tests/test_eval_tier.py +++ b/param_decomp/tests/test_eval_tier.py @@ -29,9 +29,11 @@ FAST_METRICS = { "CEandKLLossesConfig", "CIMaskedAttnPatternsReconLossConfig", + "CIMaskedReconLossConfig", "CI_L0Config", "PGDReconLossConfig", "StochasticAttnPatternsReconLossConfig", + "UnmaskedNoDeltaReconLossConfig", } SLOW_METRICS = { "ArithmeticCIGridConfig", @@ -42,7 +44,9 @@ "IdentityCIErrorConfig", "PermutedCIPlotsConfig", "StochasticHiddenActsReconLossConfig", + "TwoStreamCIMeanPerComponentConfig", "UVPlotsConfig", + "WeightMagnitudeConfig", "WellTemperednessConfig", }