Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion model.py
Original file line number Diff line number Diff line change
Expand Up @@ -379,7 +379,7 @@ def add_embedding_gaussian_noise(self, embeddings, iter_num=None):
return embeddings + noise
return embeddings

def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict=None, dataset_idx=None, loss_fn=None):
def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict=None, dataset_idx=None, loss_fn=None, return_hidden=False):
if token_dict is not None:
token_list = list(token_dict.values())
# If target_dict is None (typical for inference), set target_list = None
Expand Down Expand Up @@ -531,6 +531,8 @@ def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict
logits = [logit[:, [-1], :] for logit in logits]
losses = None

if return_hidden:
raise ValueError("return_hidden is only supported for single-context forward passes.")
return logits, losses

else:
Expand Down Expand Up @@ -633,6 +635,8 @@ def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict

loss = None

if return_hidden:
return logits, loss, x
return logits, loss
# ------------------------------------------------------------------
# LATENT-CHAINING
Expand Down
6 changes: 6 additions & 0 deletions optimization_and_search/run_experiments.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,8 @@
"ln_f_cosine_95",
"rankme",
"areq",
"target_lm_head_kl_val",
"target_lm_head_kl_train",
"zeus_total_energy_j",
"zeus_total_time_s",
"zeus_avg_power_w",
Expand Down Expand Up @@ -582,6 +584,8 @@ def read_metrics(out_dir: str) -> dict:
"ln_f_cosine_95",
"rankme",
"areq",
"target_lm_head_kl_val",
"target_lm_head_kl_train",
]
casts = [
float,
Expand All @@ -606,6 +610,8 @@ def read_metrics(out_dir: str) -> dict:
float,
float,
float,
float,
float,
]

if len(base_metric_keys) != len(casts):
Expand Down
36 changes: 29 additions & 7 deletions optimization_and_search/run_from_yaml.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,29 @@
METRICS_FILENAME = "best_val_loss_and_iter.txt"
METRIC_KEYS = [
"best_val_loss",
"best_val_iter",
"best_val_iter",
"best_tokens",
"num_params",
"better_than_chance",
Comment on lines 25 to +30
"btc_per_param",
"peak_torch_allocated_mb",
"peak_torch_reserved_mb",
"peak_process_gpu_mb",
"iter_latency_avg",
"zeus_best_train_step_energy_j",
"avg_top1_prob",
"avg_top1_correct",
"avg_target_rank",
"avg_target_left_prob",
"avg_target_prob",
"target_rank_95",
"left_prob_95",
"avg_ln_f_cosine",
"ln_f_cosine_95",
"rankme",
"areq",
"target_lm_head_kl_val",
"target_lm_head_kl_train",
]

def _parse_override_args(arg_list: list[str] | None) -> dict:
Expand Down Expand Up @@ -78,12 +98,14 @@ def read_metrics(out_dir: str) -> dict:
line = path.read_text().strip()
parts = [p.strip() for p in line.split(',')]

# Take only the first 4 values and cast them appropriately
if len(parts) < len(METRIC_KEYS):
raise ValueError(f"Expected at least {len(METRIC_KEYS)} metrics, got {len(parts)}")

casts = [float, int, int, int]
return {k: typ(v) for k, typ, v in zip(METRIC_KEYS, casts, parts[:len(METRIC_KEYS)])}
if len(parts) < 4:
raise ValueError(f"Expected at least 4 metrics, got {len(parts)}")

casts = [float, int, int, int] + [float] * (len(METRIC_KEYS) - 4)
metrics = {}
for key, typ, value in zip(METRIC_KEYS, casts, parts[:len(METRIC_KEYS)]):
metrics[key] = float("nan") if value == "" else typ(value)
return metrics
Comment on lines +104 to +108


def completed_runs(log_file: Path) -> set[str]:
Expand Down
4 changes: 4 additions & 0 deletions run_exploration_monitor.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,8 @@ def on_mount(self) -> None:
"ln_f_cosine_95",
"rankme",
"areq",
"target_lm_head_kl_val",
"target_lm_head_kl_train",
"zeus_total_energy_j",
"zeus_total_time_s",
"zeus_avg_power_w",
Expand Down Expand Up @@ -321,6 +323,8 @@ def get_cell(self, entry: Dict, col_name: str):
"ln_f_cosine_95",
"rankme",
"areq",
"target_lm_head_kl_val",
"target_lm_head_kl_train",
"zeus_total_energy_j",
"zeus_total_time_s",
"zeus_avg_power_w",
Expand Down
Loading