From 6659b4ba078b94cbf4fe6a648df2fa50f748e64c Mon Sep 17 00:00:00 2001 From: armanbolatov Date: Tue, 25 Nov 2025 22:57:30 +0000 Subject: [PATCH] removed taia, added adamuon, rmsspectral (sania), adam sania --- extract_wandb_results.py | 660 ++++++------------ requirements.txt | 2 +- .../glue/{madam => precond}/deberta_full.sh | 1 - .../glue/precond/deberta_full_adam_sania.sh | 24 + .../glue/precond/deberta_full_adamuon_ns.sh | 25 + .../glue/precond/deberta_full_adamuon_pe.sh | 25 + .../{madam => precond}/deberta_full_adamw.sh | 1 - .../{madam => precond}/deberta_full_madam.sh | 1 - scripts/glue/precond/deberta_full_madam_pe.sh | 26 + .../{madam => precond}/deberta_full_muon.sh | 1 - scripts/glue/precond/deberta_full_muon_pe.sh | 25 + .../precond/deberta_full_rmsspectral_ns.sh | 26 + .../deberta_full_rmsspectral_sania_ns.sh | 26 + .../deberta_full_rmsspectral_sania_pe.sh | 26 + .../glue/precond/deberta_lora_adam_sania.sh | 27 + .../glue/precond/deberta_lora_adamuon_ns.sh | 28 + .../{madam => precond}/deberta_lora_adamw.sh | 1 - .../{madam => precond}/deberta_lora_madam.sh | 1 - scripts/glue/precond/deberta_lora_madam_pe.sh | 31 + .../{madam => precond}/deberta_lora_muon.sh | 1 - scripts/glue/precond/deberta_lora_muon_pe.sh | 28 + .../{madam => precond}/deberta_lora_rite.sh | 1 - scripts/glue/precond/deberta_lora_rite2.sh | 27 + .../precond/deberta_lora_rmsspectral_ns.sh | 29 + .../deberta_lora_rmsspectral_sania_ns.sh | 29 + scripts/glue/{madam => precond}/run_all.sh | 0 src/config.py | 26 +- src/optimizers/adam_sania.py | 82 +++ src/optimizers/adamuon.py | 168 +++++ src/optimizers/main.py | 49 +- src/optimizers/muon.py | 45 +- src/optimizers/rmsspectral.py | 180 +++++ src/optimizers/taia.py | 348 --------- src/utils.py | 2 +- wandb_results/glue_results.csv | 91 +++ wandb_results/glue_table.tex | 26 + 36 files changed, 1270 insertions(+), 819 deletions(-) rename scripts/glue/{madam => precond}/deberta_full.sh (96%) create mode 100644 scripts/glue/precond/deberta_full_adam_sania.sh create mode 100644 scripts/glue/precond/deberta_full_adamuon_ns.sh create mode 100644 scripts/glue/precond/deberta_full_adamuon_pe.sh rename scripts/glue/{madam => precond}/deberta_full_adamw.sh (96%) rename scripts/glue/{madam => precond}/deberta_full_madam.sh (96%) create mode 100644 scripts/glue/precond/deberta_full_madam_pe.sh rename scripts/glue/{madam => precond}/deberta_full_muon.sh (96%) create mode 100644 scripts/glue/precond/deberta_full_muon_pe.sh create mode 100644 scripts/glue/precond/deberta_full_rmsspectral_ns.sh create mode 100644 scripts/glue/precond/deberta_full_rmsspectral_sania_ns.sh create mode 100644 scripts/glue/precond/deberta_full_rmsspectral_sania_pe.sh create mode 100644 scripts/glue/precond/deberta_lora_adam_sania.sh create mode 100644 scripts/glue/precond/deberta_lora_adamuon_ns.sh rename scripts/glue/{madam => precond}/deberta_lora_adamw.sh (96%) rename scripts/glue/{madam => precond}/deberta_lora_madam.sh (96%) create mode 100644 scripts/glue/precond/deberta_lora_madam_pe.sh rename scripts/glue/{madam => precond}/deberta_lora_muon.sh (96%) create mode 100644 scripts/glue/precond/deberta_lora_muon_pe.sh rename scripts/glue/{madam => precond}/deberta_lora_rite.sh (96%) create mode 100755 scripts/glue/precond/deberta_lora_rite2.sh create mode 100644 scripts/glue/precond/deberta_lora_rmsspectral_ns.sh create mode 100644 scripts/glue/precond/deberta_lora_rmsspectral_sania_ns.sh rename scripts/glue/{madam => precond}/run_all.sh (100%) create mode 100644 src/optimizers/adam_sania.py create mode 100644 src/optimizers/adamuon.py create mode 100644 src/optimizers/rmsspectral.py delete mode 100644 src/optimizers/taia.py create mode 100644 wandb_results/glue_results.csv create mode 100644 wandb_results/glue_table.tex diff --git a/extract_wandb_results.py b/extract_wandb_results.py index 0a55e86..dd582bf 100644 --- a/extract_wandb_results.py +++ b/extract_wandb_results.py @@ -1,490 +1,274 @@ import os -import re -import ast import csv -import sys -from typing import Dict, Any, Optional, Tuple, List, Set +from typing import Dict, Any -import yaml +try: + import wandb + WANDB_AVAILABLE = True +except ImportError: + WANDB_AVAILABLE = False + print("Error: wandb not installed. Run: pip install wandb") + exit(1) -WANDB_DIR = os.path.join(os.path.dirname(__file__), "wandb") +# Configuration +WANDB_PROJECT = "QUASI_DESCENT" +WANDB_ENTITY = "steeldream" OUT_DIR = os.path.join(os.path.dirname(__file__), "wandb_results") -# Default ft_strategy filter - can be overridden via command line -DEFAULT_FT_STRATEGY = "Full" - - -DATASET_METRIC_KEY = { - "cola": "eval_matthews_correlation", - "mnli": "eval_accuracy", - "mrpc": "eval_accuracy", - "qnli": "eval_accuracy", - "qqp": "eval_accuracy", - "rte": "eval_accuracy", - "sst2": "eval_accuracy", - "stsb": "eval_combined_score", - "wnli": "eval_accuracy", +OPTIMIZER_WHITELIST = { + "rmsspectral", + "rmsspectral_sania", + "adamuon", + "muon", + "adamw", } - -def read_yaml(path: str) -> Optional[Dict[str, Any]]: - try: - with open(path, "r") as f: - return yaml.safe_load(f) - except Exception: - return None - - -def parse_config(config_path: str) -> Optional[Dict[str, Any]]: - cfg = read_yaml(config_path) - if not cfg or not isinstance(cfg, dict): - return None - - def get_val(key: str) -> Optional[Any]: - node = cfg.get(key) - if isinstance(node, dict) and "value" in node: - return node.get("value") - return None - - dataset = get_val("dataset") or get_val("finetuning_task") - optimizer = get_val("optimizer") or get_val("optim") - ft_strategy = get_val("ft_strategy") - lr = get_val("lr") or get_val("learning_rate") - - # Normalize - if isinstance(optimizer, str): - optimizer = optimizer.strip().lower() - if isinstance(dataset, str): - dataset = dataset.strip().lower() - if isinstance(ft_strategy, str): - ft_strategy = ft_strategy.strip() - - return { - "dataset": dataset, - "optimizer": optimizer, - "ft_strategy": ft_strategy, - "lr": lr, - } - - -EVAL_LINE_PATTERN = re.compile(r"^\{.*\}$") - - -def extract_best_metric_from_output(log_path: str, target_key: str) -> Optional[Tuple[float, Dict[str, Any]]]: - best_val: Optional[float] = None - best_record: Optional[Dict[str, Any]] = None - try: - with open(log_path, "r") as f: - for line in f: - s = line.strip() - if "'eval_" not in s and '"eval_' not in s: - continue - if not EVAL_LINE_PATTERN.match(s): - continue - try: - rec = ast.literal_eval(s) - if not isinstance(rec, dict): - continue - except Exception: - continue - if target_key in rec and isinstance(rec[target_key], (int, float)): - val = float(rec[target_key]) - if best_val is None or val > best_val: - best_val = val - best_record = rec - except FileNotFoundError: - return None - except Exception: - return None - - if best_val is None: - return None - return best_val, (best_record or {}) - - -def scan_runs(target_ft_strategy: str = DEFAULT_FT_STRATEGY) -> Dict[str, Dict[str, Dict[str, Any]]]: - # results[dataset][optimizer] = {"best_value": float|"none", "best_lr": str|float|None, "run_id": str|None, "metric_key": str} - results: Dict[str, Dict[str, Dict[str, Any]]] = {} - - if not os.path.isdir(WANDB_DIR): - return results - - for name in os.listdir(WANDB_DIR): - run_dir = os.path.join(WANDB_DIR, name) - files_dir = os.path.join(run_dir, "files") - if not (name.startswith("run-") and os.path.isdir(files_dir)): - continue - - config_path = os.path.join(files_dir, "config.yaml") - output_log = os.path.join(files_dir, "output.log") - - cfg = parse_config(config_path) - if not cfg: - continue - - dataset = cfg.get("dataset") - optimizer = cfg.get("optimizer") - ft_strategy = cfg.get("ft_strategy") - lr = cfg.get("lr") - - if ft_strategy != target_ft_strategy: - continue - if dataset not in DATASET_METRIC_KEY: - continue - if not optimizer: - optimizer = "unknown" - - metric_key = DATASET_METRIC_KEY[dataset] - best_info = extract_best_metric_from_output(output_log, metric_key) - - if dataset not in results: - results[dataset] = {} - - if optimizer not in results[dataset]: - results[dataset][optimizer] = { - "metric_key": metric_key, - "best_value": None, - "best_lr": None, - "run_id": None, - } - - if best_info is None: - continue - - best_val, record = best_info - prev = results[dataset][optimizer]["best_value"] - if prev is None or (isinstance(prev, (int, float)) and best_val > float(prev)): - results[dataset][optimizer]["best_value"] = best_val - results[dataset][optimizer]["best_lr"] = lr - results[dataset][optimizer]["run_id"] = name - - # Replace missing with "none" - for ds in list(results.keys()): - for opt in list(results[ds].keys()): - if results[ds][opt]["best_value"] is None: - results[ds][opt]["best_value"] = "none" - return results +DATASETS = ["cola", "mnli", "mrpc", "qnli", "qqp", "rte", "sst2", "stsb", "wnli"] + +# Metric keys in W&B API format +METRIC_MAP = { + "cola": "eval/matthews_correlation", + "mnli": "eval/accuracy", + "mrpc": "eval/accuracy", + "qnli": "eval/accuracy", + "qqp": "eval/accuracy", + "rte": "eval/accuracy", + "sst2": "eval/accuracy", + "stsb": "eval/combined_score", + "wnli": "eval/accuracy", +} -def scan_runs_both_strategies() -> Dict[str, Dict[str, Dict[str, Dict[str, Any]]]]: - # results[ft_strategy][dataset][optimizer] = {"best_value": float|"none", "best_lr": str|float|None, "run_id": str|None, "metric_key": str} - results: Dict[str, Dict[str, Dict[str, Dict[str, Any]]]] = {} - strategies = ["LoRA", "Full"] +def fetch_results() -> Dict[str, Dict[str, Dict[str, Any]]]: + """Fetch results from W&B API for both LoRA and Full strategies.""" + api = wandb.Api() + project_path = f"{WANDB_ENTITY}/{WANDB_PROJECT}" - for strategy in strategies: - results[strategy] = scan_runs(strategy) + print(f"Fetching runs from W&B project: {project_path}") + runs = api.runs(project_path) + + # Structure: results[ft_strategy][dataset][optimizer] = {best_value, lr, run_id} + results = {"LoRA": {}, "Full": {}} + + for strategy in ["LoRA", "Full"]: + for dataset in DATASETS: + results[strategy][dataset] = {} + + total = 0 + processed = 0 + + for run in runs: + total += 1 + try: + config = run.config + summary = run.summary._json_dict + + # Extract config + dataset = (config.get("dataset") or "").strip().lower() + optimizer = (config.get("optimizer") or "").strip().lower().replace('-', '_') + ft_strategy = (config.get("ft_strategy") or "").strip() + lr = config.get("lr") or config.get("learning_rate") + + # Filter + if ft_strategy not in ["LoRA", "Full"]: + continue + if dataset not in DATASETS: + continue + if optimizer not in OPTIMIZER_WHITELIST: + continue + + # Get metric + metric_key = METRIC_MAP[dataset] + best_val = summary.get(metric_key) + + if not isinstance(best_val, (int, float)): + continue + + # Update best for this optimizer/dataset/strategy + current = results[ft_strategy][dataset].get(optimizer) + if current is None or best_val > current.get("best_value", 0): + results[ft_strategy][dataset][optimizer] = { + "best_value": best_val, + "lr": lr, + "run_id": run.id, + } + + processed += 1 + + except Exception: + continue + print(f"Processed {processed}/{total} runs") return results -def write_csvs(results: Dict[str, Dict[str, Dict[str, Any]]], ft_strategy: str) -> None: +def write_csv(results: Dict[str, Dict[str, Dict[str, Any]]]) -> None: + """Write single combined CSV with all results.""" os.makedirs(OUT_DIR, exist_ok=True) - ft_suffix = ft_strategy.lower() - for dataset, by_opt in results.items(): - out_path = os.path.join(OUT_DIR, f"glue_{dataset}_{ft_suffix}.csv") - rows = [] - metric_key = DATASET_METRIC_KEY.get(dataset, "") - for optimizer, info in sorted(by_opt.items()): - rows.append({ - "optimizer": optimizer, - "metric": metric_key, - "best_value": info.get("best_value", "none"), - "best_lr": info.get("best_lr"), - "run_id": info.get("run_id"), - }) - with open(out_path, "w", newline="") as f: - writer = csv.DictWriter(f, fieldnames=["optimizer", "metric", "best_value", "best_lr", "run_id"]) - writer.writeheader() - for r in rows: - writer.writerow(r) - -def write_combined_csv(combined_results: Dict[str, Dict[str, Dict[str, Dict[str, Any]]]]) -> None: - os.makedirs(OUT_DIR, exist_ok=True) - out_path = os.path.join(OUT_DIR, "glue_combined_results.csv") + out_path = os.path.join(OUT_DIR, "glue_results.csv") rows = [] for ft_strategy in ["LoRA", "Full"]: - if ft_strategy not in combined_results: - continue - for dataset, by_opt in combined_results[ft_strategy].items(): - metric_key = DATASET_METRIC_KEY.get(dataset, "") - for optimizer, info in sorted(by_opt.items()): - rows.append({ - "ft_strategy": ft_strategy, - "dataset": dataset, - "optimizer": optimizer, - "metric": metric_key, - "best_value": info.get("best_value", "none"), - "best_lr": info.get("best_lr"), - "run_id": info.get("run_id"), - }) + for dataset in DATASETS: + for optimizer in sorted(OPTIMIZER_WHITELIST): + result = results[ft_strategy][dataset].get(optimizer) + if result: + rows.append({ + "ft_strategy": ft_strategy, + "dataset": dataset, + "optimizer": optimizer, + "best_value": f"{result['best_value']:.4f}", + "lr": result["lr"], + "run_id": result["run_id"], + }) + else: + rows.append({ + "ft_strategy": ft_strategy, + "dataset": dataset, + "optimizer": optimizer, + "best_value": "n/a", + "lr": "", + "run_id": "", + }) with open(out_path, "w", newline="") as f: - writer = csv.DictWriter(f, fieldnames=["ft_strategy", "dataset", "optimizer", "metric", "best_value", "best_lr", "run_id"]) + writer = csv.DictWriter( + f, + fieldnames=["ft_strategy", "dataset", "optimizer", "best_value", "lr", "run_id"] + ) writer.writeheader() - for r in rows: - writer.writerow(r) - - -def _format_value(v: Any) -> str: - if v is None: - return "n/a" - if isinstance(v, str): - return v if v.strip() else "n/a" - try: - x = float(v) - return f"{x:.4f}" - except Exception: - return "n/a" - - -def write_latex_table_from_results(results: Dict[str, Dict[str, Dict[str, Any]]], ft_strategy: str) -> None: - os.makedirs(OUT_DIR, exist_ok=True) - optimizers: List[str] = [] - seen: Set[str] = set() - for by_opt in results.values(): - for opt in by_opt.keys(): - if opt not in seen: - seen.add(opt) - optimizers.append(opt) - optimizers.sort() - - ft_suffix = ft_strategy.lower() - table_path = os.path.join(OUT_DIR, f"glue_results_table_{ft_suffix}.tex") - - header_opt_cols = " & ".join(opt.upper() for opt in optimizers) if optimizers else "" - - lines: List[str] = [] - lines.append("% Auto-generated by extract_wandb_results.py") - lines.append("\\begin{table}[t]") - lines.append("\\centering") - col_spec = "l l" + (" " + "c" * len(optimizers) if optimizers else "") - lines.append(f"\\begin{{tabular}}{{{col_spec}}}") - lines.append("\\toprule") - if optimizers: - lines.append(f"Dataset & Metric & {header_opt_cols} \\\ ") - else: - lines.append("Dataset & Metric \\\ ") - lines.append("\\midrule") - - glue_order = [ - "cola", "mnli", "mrpc", "qnli", "qqp", "rte", "sst2", "stsb", "wnli" - ] - - def pretty_metric_key(ds: str) -> str: - raw = DATASET_METRIC_KEY.get(ds, "") - mapping = { - "eval_matthews_correlation": "Matthews", - "eval_accuracy": "Accuracy", - "eval_combined_score": "Combined", - } - return mapping.get(raw, raw) - - for ds in glue_order: - if ds not in results: - metric_name = pretty_metric_key(ds) - vals = ["n/a" for _ in optimizers] - else: - metric_name = pretty_metric_key(ds) - row = results[ds] - vals = [] - for opt in optimizers: - best = row.get(opt, {}).get("best_value") - if isinstance(best, str) and best == "none": - vals.append("n/a") - else: - vals.append(_format_value(best)) - if optimizers: - lines.append(f"{ds.upper()} & {metric_name} & " + " & ".join(vals) + " \\") - else: - lines.append(f"{ds.upper()} & {metric_name} \\") - - lines.append("\\bottomrule") - lines.append("\\end{tabular}") - lines.append(f"\\caption{{GLUE {ft_strategy} fine-tuning: best validation metrics per optimizer (higher is better). Missing entries are denoted as n/a.}}") - lines.append("\\label{tab:glue_results}") - lines.append("\\end{table}") - - with open(table_path, "w") as f: - f.write("\n".join(lines)) + writer.writerows(rows) + + print(f"\nWrote results to: {out_path}") -def write_combined_latex_table(combined_results: Dict[str, Dict[str, Dict[str, Dict[str, Any]]]]) -> None: +def write_latex_table(results: Dict[str, Dict[str, Dict[str, Any]]]) -> None: + """Write LaTeX table with both LoRA and Full results.""" os.makedirs(OUT_DIR, exist_ok=True) + out_path = os.path.join(OUT_DIR, "glue_table.tex") - # Collect all optimizers from both strategies - optimizers: List[str] = [] - seen: Set[str] = set() - for strategy_results in combined_results.values(): - for by_opt in strategy_results.values(): - for opt in by_opt.keys(): - if opt not in seen: - seen.add(opt) - optimizers.append(opt) - optimizers.sort() - - table_path = os.path.join(OUT_DIR, "glue_combined_results_table.tex") - - lines: List[str] = [] - lines.append("% Auto-generated by extract_wandb_results.py") - lines.append("\\begin{table}[t]") - lines.append("\\centering") - lines.append("\\scriptsize") - lines.append("\\setlength{\\tabcolsep}{3pt}") - lines.append("\\renewcommand{\\arraystretch}{1.1}") - lines.append("\\resizebox{0.95\\linewidth}{!}{%") + # Dataset order and their short names for table + datasets_short = ["cola", "mnli", "mrpc", "qnli", "qqp", "rte", "sst2", "stsb"] + opt_map = { + "adamw": "AdamW", + "muon": "Muon", + "adamuon": "AdaMuon", + "rmsspectral": "RMSSpectral", + "rmsspectral_sania": "RMSSpectral-SANIA" + } - # Create column specification - col_spec = "l l" + "c" * 9 + "|c" # 9 datasets + average column - lines.append(f"\\begin{{tabular}}{{{col_spec}}}") + lines = [ + "\\begin{table}[h!]", + "\\centering", + "\\scriptsize", + "\\setlength{\\tabcolsep}{3pt}", + "\\renewcommand{\\arraystretch}{1.1}", + "\\captionof{table}{GLUE: datasets are columns with the corresponding metric; \\texttt{ALL} is the average over tasks. Best per task in bold.}", + "\\resizebox{0.95\\linewidth}{!}{", + "\\begin{tabular}{l lcccccccc|c}", + "& & \\begin{tabular}{@{}c@{}}CoLA\\\\Matthews\\end{tabular} & \\begin{tabular}{@{}c@{}}MNLI\\\\Acc\\end{tabular} & \\begin{tabular}{@{}c@{}}MRPC\\\\Acc\\end{tabular} & \\begin{tabular}{@{}c@{}}QNLI\\\\Acc\\end{tabular} & \\begin{tabular}{@{}c@{}}QQP\\\\Acc\\end{tabular} & \\begin{tabular}{@{}c@{}}RTE\\\\Acc\\end{tabular} & \\begin{tabular}{@{}c@{}}SST-2\\\\Acc\\end{tabular} & \\begin{tabular}{@{}c@{}}STS-B\\\\Comb.\\end{tabular} & \\begin{tabular}{@{}c@{}}ALL\\\\Avg\\end{tabular} \\\\", + "\\midrule" + ] - # Header with dataset columns - lines.append(" % \\toprule") - header_line = "& & \\begin{tabular}{@{}c@{}}CoLA\\\\Matthews\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}MNLI\\\\Acc\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}MRPC\\\\Acc\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}QNLI\\\\Acc\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}QQP\\\\Acc\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}RTE\\\\Acc\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}SST-2\\\\Acc\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}STS-B\\\\Comb.\\end{tabular}" - header_line += " & \\begin{tabular}{@{}c@{}}ALL\\\\Avg\\end{tabular} \\\\" - lines.append(header_line) - lines.append("\\midrule") - - glue_order = ["cola", "mnli", "mrpc", "qnli", "qqp", "rte", "sst2", "stsb"] - - # Process each strategy - for strategy_idx, ft_strategy in enumerate(["LoRA", "Full"]): - if ft_strategy not in combined_results: - continue + for strategy_idx, strategy in enumerate(["LoRA", "Full"]): + # Find best values per dataset across all optimizers for this strategy + best_vals = {} + for dataset in datasets_short: + best_val = None + for optimizer in OPTIMIZER_WHITELIST: + result = results[strategy][dataset].get(optimizer) + if result and isinstance(result["best_value"], (int, float)): + if best_val is None or result["best_value"] > best_val: + best_val = result["best_value"] + best_vals[dataset] = best_val - strategy_results = combined_results[ft_strategy] + num_opts = len(OPTIMIZER_WHITELIST) - # Add strategy rows - for opt_idx, optimizer in enumerate(optimizers): - if opt_idx == 0: - # First row for this strategy - use multirow - lines.append(f"\\multirow{{{len(optimizers)}}}{{*}}{{{ft_strategy}}}") - else: - lines.append("") + for opt_idx, optimizer in enumerate(sorted(OPTIMIZER_WHITELIST)): + opt_name = opt_map.get(optimizer, optimizer.capitalize()) - # Optimizer name - opt_name = f"\\texttt{{{optimizer}}}" - - # Collect values for each dataset + # Values for each dataset values = [] - valid_values = [] # For computing average + valid_values = [] - for ds in glue_order: - if ds in strategy_results and optimizer in strategy_results[ds]: - best = strategy_results[ds][optimizer].get("best_value") - if isinstance(best, str) and best == "none": - values.append("n/a") + for dataset in datasets_short: + result = results[strategy][dataset].get(optimizer) + if result and isinstance(result["best_value"], (int, float)): + val = result["best_value"] + valid_values.append(val) + # Bold if best in this dataset for this strategy + if best_vals[dataset] and abs(val - best_vals[dataset]) < 1e-6: + values.append(f"\\textbf{{{val:.4f}}}") else: - formatted_val = _format_value(best) - values.append(formatted_val) - try: - valid_values.append(float(best)) - except (ValueError, TypeError): - pass + values.append(f"{val:.4f}") else: - values.append("n/a") + values.append("NaN") # Calculate average if valid_values: - avg_val = sum(valid_values) / len(valid_values) - avg_formatted = f"{avg_val:.4f}" - else: - avg_formatted = "n/a" - - # Find best values for bolding - best_in_strategy = {} - for ds in glue_order: - best_val = None - for opt in optimizers: - if ds in strategy_results and opt in strategy_results[ds]: - val = strategy_results[ds][opt].get("best_value") - if isinstance(val, (int, float)): - if best_val is None or val > best_val: - best_val = val - best_in_strategy[ds] = opt - - # Format values with bold for best - formatted_values = [] - for i, (ds, val) in enumerate(zip(glue_order, values)): - if val != "n/a" and ds in best_in_strategy and best_in_strategy[ds] == optimizer: - formatted_values.append(f"\\textbf{{{val}}}") + avg = sum(valid_values) / len(valid_values) + + # Check if this is best average for this strategy + best_avg = None + for opt_check in OPTIMIZER_WHITELIST: + opt_vals = [] + for ds in datasets_short: + res = results[strategy][ds].get(opt_check) + if res and isinstance(res["best_value"], (int, float)): + opt_vals.append(res["best_value"]) + if opt_vals: + opt_avg = sum(opt_vals) / len(opt_vals) + if best_avg is None or opt_avg > best_avg: + best_avg = opt_avg + + if best_avg and abs(avg - best_avg) < 1e-6: + avg_str = f"\\textbf{{{avg:.4f}}}" else: - formatted_values.append(val) - - # Check if this optimizer has best average - best_avg_opt = None - best_avg_val = None - for opt in optimizers: - opt_valid_values = [] - for ds in glue_order: - if ds in strategy_results and opt in strategy_results[ds]: - val = strategy_results[ds][opt].get("best_value") - if isinstance(val, (int, float)): - opt_valid_values.append(val) - if opt_valid_values: - opt_avg = sum(opt_valid_values) / len(opt_valid_values) - if best_avg_val is None or opt_avg > best_avg_val: - best_avg_val = opt_avg - best_avg_opt = opt + avg_str = f"{avg:.4f}" + else: + avg_str = "NaN" - if avg_formatted != "n/a" and best_avg_opt == optimizer: - avg_formatted = f"\\textbf{{{avg_formatted}}}" + # Add multirow for first optimizer of each strategy + if opt_idx == 0: + prefix = f"\\multirow{{{num_opts}}}{{*}}{{{strategy}}}" + else: + prefix = "" - # Create the row - row = f"& {opt_name} & " + " & ".join(formatted_values) + f" & {avg_formatted} \\\\" + row = f"{prefix} & \\texttt{{{opt_name}}} & {' & '.join(values)} & {avg_str} \\\\" lines.append(row) - # Add separator between strategies - if strategy_idx == 0: # After LoRA, before Full + # Add midrule between strategies (but not after last) + if strategy_idx == 0: lines.append("\\midrule") - - lines.append("\\bottomrule") - lines.append("\\end{tabular}%") - lines.append("}") - lines.append("\\caption{GLUE (LoRA and Full fine-tuning): datasets are columns with metric under the dataset name; \\texttt{ALL} is the average over tasks.}") - lines.append("\\label{tab:glue_results_transposed}") - lines.append("\\end{table}") - - with open(table_path, "w") as f: + + lines.extend([ + "\\bottomrule", + "\\end{tabular}", + "}", + "\\label{tab:glue_results_transposed}", + "\\end{table}" + ]) + + with open(out_path, "w") as f: f.write("\n".join(lines)) + + print(f"Wrote LaTeX table to: {out_path}") def main() -> None: - # Check for command line argument to specify ft_strategy - if len(sys.argv) > 1 and sys.argv[1] != "combined": - # Single strategy mode (backward compatibility) - ft_strategy = sys.argv[1] - print(f"Extracting results for ft_strategy: {ft_strategy}") - - results = scan_runs(ft_strategy) - if not results: - print(f"No wandb results found for ft_strategy='{ft_strategy}' or wandb directory missing.") - return - write_csvs(results, ft_strategy) - write_latex_table_from_results(results, ft_strategy) - print(f"Wrote CSVs to: {OUT_DIR}") - print(f"Wrote LaTeX table to: {os.path.join(OUT_DIR, f'glue_results_table_{ft_strategy.lower()}.tex')}") - else: - # Combined mode (default) - print("Extracting results for both LoRA and Full strategies...") - - combined_results = scan_runs_both_strategies() - if not combined_results or not any(combined_results.values()): - print("No wandb results found for either strategy or wandb directory missing.") - return - - write_combined_csv(combined_results) - write_combined_latex_table(combined_results) - print(f"Wrote combined CSV to: {os.path.join(OUT_DIR, 'glue_combined_results.csv')}") - print(f"Wrote combined LaTeX table to: {os.path.join(OUT_DIR, 'glue_combined_results_table.tex')}") + if not WANDB_AVAILABLE: + return + + results = fetch_results() + write_csv(results) + write_latex_table(results) + + # Print summary + print("\nSummary:") + for strategy in ["LoRA", "Full"]: + found = sum(1 for ds in results[strategy].values() for opt in ds.keys()) + print(f" {strategy}: {found} optimizer/dataset combinations") if __name__ == "__main__": diff --git a/requirements.txt b/requirements.txt index 7548bff..bf28d84 100644 --- a/requirements.txt +++ b/requirements.txt @@ -7,7 +7,7 @@ peft==0.16.0 datasets==2.19.0 # Data science -numpy==1.24.4 +numpy==1.26.4 pandas==2.2.3 matplotlib==3.9.3 seaborn==0.13.2 diff --git a/scripts/glue/madam/deberta_full.sh b/scripts/glue/precond/deberta_full.sh similarity index 96% rename from scripts/glue/madam/deberta_full.sh rename to scripts/glue/precond/deberta_full.sh index c64ecd5..a09c8c4 100755 --- a/scripts/glue/madam/deberta_full.sh +++ b/scripts/glue/precond/deberta_full.sh @@ -20,6 +20,5 @@ do --eval_strategy epoch \ --save_strategy no \ --ft_strategy Full \ - --dtype bfloat16 \ --wandb done diff --git a/scripts/glue/precond/deberta_full_adam_sania.sh b/scripts/glue/precond/deberta_full_adam_sania.sh new file mode 100644 index 0000000..59aa4a7 --- /dev/null +++ b/scripts/glue/precond/deberta_full_adam_sania.sh @@ -0,0 +1,24 @@ +clear + +# datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +datasets=(cola rte) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer adam_sania \ + --batch_size 16 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --wandb + done + done diff --git a/scripts/glue/precond/deberta_full_adamuon_ns.sh b/scripts/glue/precond/deberta_full_adamuon_ns.sh new file mode 100644 index 0000000..90ed130 --- /dev/null +++ b/scripts/glue/precond/deberta_full_adamuon_ns.sh @@ -0,0 +1,25 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer adamuon \ + --ns_steps 6 \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_full_adamuon_pe.sh b/scripts/glue/precond/deberta_full_adamuon_pe.sh new file mode 100644 index 0000000..f6750a2 --- /dev/null +++ b/scripts/glue/precond/deberta_full_adamuon_pe.sh @@ -0,0 +1,25 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=1 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer adamuon \ + --ns_steps 6 \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/madam/deberta_full_adamw.sh b/scripts/glue/precond/deberta_full_adamw.sh similarity index 96% rename from scripts/glue/madam/deberta_full_adamw.sh rename to scripts/glue/precond/deberta_full_adamw.sh index a4b9bb0..95cf73f 100644 --- a/scripts/glue/madam/deberta_full_adamw.sh +++ b/scripts/glue/precond/deberta_full_adamw.sh @@ -19,7 +19,6 @@ for dataset in "${datasets[@]}"; do --eval_strategy epoch \ --save_strategy no \ --ft_strategy Full \ - --dtype bfloat16 \ --wandb done done \ No newline at end of file diff --git a/scripts/glue/madam/deberta_full_madam.sh b/scripts/glue/precond/deberta_full_madam.sh similarity index 96% rename from scripts/glue/madam/deberta_full_madam.sh rename to scripts/glue/precond/deberta_full_madam.sh index 0f4c408..ff5d233 100644 --- a/scripts/glue/madam/deberta_full_madam.sh +++ b/scripts/glue/precond/deberta_full_madam.sh @@ -22,7 +22,6 @@ for dataset in "${datasets[@]}"; do --eval_strategy epoch \ --save_strategy no \ --ft_strategy Full \ - --dtype bfloat16 \ --wandb done done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_full_madam_pe.sh b/scripts/glue/precond/deberta_full_madam_pe.sh new file mode 100644 index 0000000..fee1db8 --- /dev/null +++ b/scripts/glue/precond/deberta_full_madam_pe.sh @@ -0,0 +1,26 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=1 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer rmsspectral \ + --rms_power 0.25 \ + --ns_steps 6 \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/madam/deberta_full_muon.sh b/scripts/glue/precond/deberta_full_muon.sh similarity index 96% rename from scripts/glue/madam/deberta_full_muon.sh rename to scripts/glue/precond/deberta_full_muon.sh index 9b8b34b..7d45c78 100644 --- a/scripts/glue/madam/deberta_full_muon.sh +++ b/scripts/glue/precond/deberta_full_muon.sh @@ -19,7 +19,6 @@ for dataset in "${datasets[@]}"; do --eval_strategy epoch \ --save_strategy no \ --ft_strategy Full \ - --dtype bfloat16 \ --wandb done done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_full_muon_pe.sh b/scripts/glue/precond/deberta_full_muon_pe.sh new file mode 100644 index 0000000..3446e13 --- /dev/null +++ b/scripts/glue/precond/deberta_full_muon_pe.sh @@ -0,0 +1,25 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +# datasets=(cola rte) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=2 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer muon \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --orth_algo ns \ + --wandb + done + done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_full_rmsspectral_ns.sh b/scripts/glue/precond/deberta_full_rmsspectral_ns.sh new file mode 100644 index 0000000..c41cc86 --- /dev/null +++ b/scripts/glue/precond/deberta_full_rmsspectral_ns.sh @@ -0,0 +1,26 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer rmsspectral \ + --rms_power 0.25 \ + --ns_steps 6 \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_full_rmsspectral_sania_ns.sh b/scripts/glue/precond/deberta_full_rmsspectral_sania_ns.sh new file mode 100644 index 0000000..f121758 --- /dev/null +++ b/scripts/glue/precond/deberta_full_rmsspectral_sania_ns.sh @@ -0,0 +1,26 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer rmsspectral_sania \ + --rms_power 0.5 \ + --ns_steps 6 \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_full_rmsspectral_sania_pe.sh b/scripts/glue/precond/deberta_full_rmsspectral_sania_pe.sh new file mode 100644 index 0000000..11db3f9 --- /dev/null +++ b/scripts/glue/precond/deberta_full_rmsspectral_sania_pe.sh @@ -0,0 +1,26 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=1 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer rmsspectral_sania \ + --rms_power 0.5 \ + --ns_steps 6 \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy Full \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_lora_adam_sania.sh b/scripts/glue/precond/deberta_lora_adam_sania.sh new file mode 100644 index 0000000..0fb890c --- /dev/null +++ b/scripts/glue/precond/deberta_lora_adam_sania.sh @@ -0,0 +1,27 @@ +clear + +# datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +datasets=(cola rte) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer adam_sania \ + --batch_size 16 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy LoRA \ + --lora_r 4 \ + --lora_alpha 32 \ + --lora_dropout 0.05 \ + --wandb + done + done diff --git a/scripts/glue/precond/deberta_lora_adamuon_ns.sh b/scripts/glue/precond/deberta_lora_adamuon_ns.sh new file mode 100644 index 0000000..369bb35 --- /dev/null +++ b/scripts/glue/precond/deberta_lora_adamuon_ns.sh @@ -0,0 +1,28 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5 1e-4) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer adamuon \ + --ns_steps 6 \ + --batch_size 16 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy LoRA \ + --lora_r 4 \ + --lora_alpha 32 \ + --lora_dropout 0.1 \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/madam/deberta_lora_adamw.sh b/scripts/glue/precond/deberta_lora_adamw.sh similarity index 96% rename from scripts/glue/madam/deberta_lora_adamw.sh rename to scripts/glue/precond/deberta_lora_adamw.sh index 0936467..daa5322 100644 --- a/scripts/glue/madam/deberta_lora_adamw.sh +++ b/scripts/glue/precond/deberta_lora_adamw.sh @@ -19,7 +19,6 @@ for dataset in "${datasets[@]}"; do --eval_strategy epoch \ --save_strategy no \ --ft_strategy LoRA \ - --dtype bfloat16 \ --lora_r 4 \ --lora_alpha 32 \ --lora_dropout 0.05 \ diff --git a/scripts/glue/madam/deberta_lora_madam.sh b/scripts/glue/precond/deberta_lora_madam.sh similarity index 96% rename from scripts/glue/madam/deberta_lora_madam.sh rename to scripts/glue/precond/deberta_lora_madam.sh index c3fdecb..bd8d46f 100755 --- a/scripts/glue/madam/deberta_lora_madam.sh +++ b/scripts/glue/precond/deberta_lora_madam.sh @@ -22,7 +22,6 @@ for dataset in "${datasets[@]}"; do --eval_strategy epoch \ --save_strategy no \ --ft_strategy LoRA \ - --dtype bfloat16 \ --lora_r 4 \ --lora_alpha 32 \ --lora_dropout 0.05 \ diff --git a/scripts/glue/precond/deberta_lora_madam_pe.sh b/scripts/glue/precond/deberta_lora_madam_pe.sh new file mode 100644 index 0000000..1bf0fc1 --- /dev/null +++ b/scripts/glue/precond/deberta_lora_madam_pe.sh @@ -0,0 +1,31 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +# datasets=(cola rte) +lrs=(3e-4 1e-3 2e-5 1e-4) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer taia \ + --lmo spectral \ + --precondition_type adam \ + --init eps \ + --batch_size 32 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy LoRA \ + --lora_r 4 \ + --lora_alpha 32 \ + --lora_dropout 0.1 \ + --orth_algo polar \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/madam/deberta_lora_muon.sh b/scripts/glue/precond/deberta_lora_muon.sh similarity index 96% rename from scripts/glue/madam/deberta_lora_muon.sh rename to scripts/glue/precond/deberta_lora_muon.sh index 4aacb2f..922c1a4 100644 --- a/scripts/glue/madam/deberta_lora_muon.sh +++ b/scripts/glue/precond/deberta_lora_muon.sh @@ -19,7 +19,6 @@ for dataset in "${datasets[@]}"; do --eval_strategy epoch \ --save_strategy no \ --ft_strategy LoRA \ - --dtype bfloat16 \ --lora_r 4 \ --lora_alpha 32 \ --lora_dropout 0.05 \ diff --git a/scripts/glue/precond/deberta_lora_muon_pe.sh b/scripts/glue/precond/deberta_lora_muon_pe.sh new file mode 100644 index 0000000..cc2f72c --- /dev/null +++ b/scripts/glue/precond/deberta_lora_muon_pe.sh @@ -0,0 +1,28 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +# datasets=(cola rte) +lrs=(3e-4 1e-3 2e-5 1e-4) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer muon \ + --batch_size 16 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy LoRA \ + --lora_r 4 \ + --lora_alpha 32 \ + --lora_dropout 0.1 \ + --orth_algo ns \ + --wandb + done + done \ No newline at end of file diff --git a/scripts/glue/madam/deberta_lora_rite.sh b/scripts/glue/precond/deberta_lora_rite.sh similarity index 96% rename from scripts/glue/madam/deberta_lora_rite.sh rename to scripts/glue/precond/deberta_lora_rite.sh index 2e11562..9f2b6e1 100755 --- a/scripts/glue/madam/deberta_lora_rite.sh +++ b/scripts/glue/precond/deberta_lora_rite.sh @@ -19,7 +19,6 @@ for dataset in "${datasets[@]}"; do --eval_strategy epoch \ --save_strategy no \ --ft_strategy LoRA \ - --dtype bfloat16 \ --lora_r 4 \ --lora_alpha 32 \ --lora_dropout 0.05 \ diff --git a/scripts/glue/precond/deberta_lora_rite2.sh b/scripts/glue/precond/deberta_lora_rite2.sh new file mode 100755 index 0000000..37c4ab4 --- /dev/null +++ b/scripts/glue/precond/deberta_lora_rite2.sh @@ -0,0 +1,27 @@ +clear + +datasets=(qqp rte sst2 stsb) +# datasets=(cola rte) +lrs=(3e-4 1e-3 2e-5) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer lora_rite \ + --batch_size 16 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy LoRA \ + --lora_r 4 \ + --lora_alpha 32 \ + --lora_dropout 0.05 \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_lora_rmsspectral_ns.sh b/scripts/glue/precond/deberta_lora_rmsspectral_ns.sh new file mode 100644 index 0000000..2a7d09d --- /dev/null +++ b/scripts/glue/precond/deberta_lora_rmsspectral_ns.sh @@ -0,0 +1,29 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5 1e-4) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer rmsspectral \ + --rms_power 0.25 \ + --ns_steps 6 \ + --batch_size 16 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy LoRA \ + --lora_r 4 \ + --lora_alpha 32 \ + --lora_dropout 0.1 \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/precond/deberta_lora_rmsspectral_sania_ns.sh b/scripts/glue/precond/deberta_lora_rmsspectral_sania_ns.sh new file mode 100644 index 0000000..6ec8658 --- /dev/null +++ b/scripts/glue/precond/deberta_lora_rmsspectral_sania_ns.sh @@ -0,0 +1,29 @@ +clear + +datasets=(cola mnli mrpc qnli qqp rte sst2 stsb) +lrs=(3e-4 1e-3 2e-5 1e-4) + +for dataset in "${datasets[@]}"; do + for lr in "${lrs[@]}"; do + CUDA_VISIBLE_DEVICES=0 python ./src/run_experiment.py \ + --dataset $dataset \ + --model distilbert/distilbert-base-uncased \ + --optimizer rmsspectral_sania \ + --rms_power 0.5 \ + --ns_steps 6 \ + --batch_size 16 \ + --gradient_accumulation_steps 2 \ + --lr $lr \ + --weight_decay 0.1 \ + --lr_scheduler_type linear \ + --warmup_ratio 0.1 \ + --max_train_steps 10000 \ + --eval_strategy epoch \ + --save_strategy no \ + --ft_strategy LoRA \ + --lora_r 4 \ + --lora_alpha 32 \ + --lora_dropout 0.1 \ + --wandb + done +done \ No newline at end of file diff --git a/scripts/glue/madam/run_all.sh b/scripts/glue/precond/run_all.sh similarity index 100% rename from scripts/glue/madam/run_all.sh rename to scripts/glue/precond/run_all.sh diff --git a/src/config.py b/src/config.py index f8ec662..32ed58a 100644 --- a/src/config.py +++ b/src/config.py @@ -101,7 +101,7 @@ def parse_args(): ) parser.add_argument("--eps", default=1e-8, type=float, help="Epsilon for Adam") - if args1.optimizer in ["shampoo", "sgd", "muon", "taia"]: + if args1.optimizer in ["shampoo", "sgd", "muon", "adamuon", "rmsspectral", "rmsspectral_sania"]: parser.add_argument( "--momentum", default=0.9, type=float, help="First momentum" ) @@ -119,32 +119,30 @@ def parse_args(): type=int, help="maximum dimension of preconditioner for SOAP-like algorithms", ) - if args1.optimizer in ["shampoo", "soap", "diag-hvp", "taia"]: + if args1.optimizer in ["shampoo", "soap", "diag-hvp"]: parser.add_argument( "--update_freq", default=1, type=int, help="Freqiensy to update Q for Shampoo and SOAP", ) - if args1.optimizer in ["muon", "taia"]: + if args1.optimizer in ["muon", "adamuon", "rmsspectral", "rmsspectral_sania"]: parser.add_argument( "--ns_steps", default=10, type=int, help="Number of the NS steps algo" ) parser.add_argument( "--adamw_lr", default=1e-4, type=float, help="lr for adam in " ) - if args1.optimizer == "taia": - parser.add_argument( - "--lmo", - default="frobenious", - type=str, - help="Linear minimization oracle for taia", - ) + if args1.optimizer == "muon": + parser.add_argument( + "--orth_algo", + default="ns", + type=str, + help="Orthogonalization algorithm: 'ns' (Newton–Schulz) or 'polar' (PolarExpress)", + ) + if args1.optimizer in ["rmsspectral", "rmsspectral_sania"]: parser.add_argument( - "--precondition_type", - default="norm", - type=str, - help="Type of preconditioner for taia", + "--rms_power", default=0.25, type=float, help="RMS spectral power (0.25 default, 0.5 for SANIA variant)" ) ### Problem Specific Arguments diff --git a/src/optimizers/adam_sania.py b/src/optimizers/adam_sania.py new file mode 100644 index 0000000..8b49093 --- /dev/null +++ b/src/optimizers/adam_sania.py @@ -0,0 +1,82 @@ +import torch +from typing import Iterable, Callable + + +class AdamSania(torch.optim.Optimizer): + """A very small standalone Adam implementation. + + Mirrors the structure of `Muon` (custom optimizer class) but only performs + the classic Adam update without any orthogonalization or auxiliary logic. + + Arguments: + params: Iterable of parameters to optimize. + lr: Learning rate (default: 1e-4) + betas: Coefficients used for computing running averages of gradient + and its square (default: (0.9, 0.999)) + eps: Term added to the denominator for numerical stability (default: 1e-8) + weight_decay: L2 penalty (applied in-gradient, not decoupled) (default: 0.0) + """ + + def __init__( + self, + params: Iterable[torch.nn.Parameter], + lr: float = 1e-4, + betas=(0.9, 0.999), + eps: float = 1e-8, + weight_decay: float = 0.0, + ): + if lr <= 0.0: + raise ValueError(f"Invalid learning rate: {lr}") + if eps <= 0.0: + raise ValueError(f"Invalid epsilon value: {eps}") + if not 0.0 <= betas[0] < 1.0: + raise ValueError(f"Invalid beta1 value: {betas[0]}") + if not 0.0 <= betas[1] < 1.0: + raise ValueError(f"Invalid beta2 value: {betas[1]}") + + defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay) + super().__init__(params, defaults) + + def step(self, closure: Callable = None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + + for group in self.param_groups: + beta1, beta2 = group["betas"] + eps = group["eps"] + lr = group["lr"] + wd = group["weight_decay"] + + for p in group["params"]: + if p.grad is None: + continue + g = p.grad.data + if wd != 0.0: + g = g.add(p.data, alpha=wd) + + state = self.state[p] + if len(state) == 0: + state["step"] = 0 + state["exp_avg"] = torch.zeros_like(p.data) + state["exp_avg_sq"] = torch.zeros_like(p.data) + + exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"] + state["step"] += 1 + step = state["step"] + + # Update first and second moments + exp_avg.mul_(beta1).add_(g, alpha=1 - beta1) + exp_avg_sq.mul_(beta2).addcmul_(g, g, value=1 - beta2) + + # Bias correction + bias_correction1 = 1 - beta1 ** step + bias_correction2 = 1 - beta2 ** step + denom = (exp_avg_sq.abs() / (bias_correction2)).add_(eps) + print("denom:", denom) + step_size = lr / bias_correction1 + + p.data.addcdiv_(exp_avg, denom, value=-step_size) + + return loss diff --git a/src/optimizers/adamuon.py b/src/optimizers/adamuon.py new file mode 100644 index 0000000..448b79d --- /dev/null +++ b/src/optimizers/adamuon.py @@ -0,0 +1,168 @@ +"""AdaMuon optimizer implementation (replaces Muon placeholder). + +Algorithm (per paper draft provided): +M_t = beta * M_{t-1} + G_t +O_t = NewtonSchulz(Sign(M_t)) +V_t = beta * V_{t-1} + (1 - beta) * (O_t ⊙ O_t) +O_hat = O_t / (sqrt(V_t) + eps) +γ_t = 0.2 * sqrt(m*n) / ||O_hat||_F +W_{t+1} = W_t - lr * (γ_t * O_hat + weight_decay * W_t) + +We only apply AdaMuon to 2D parameters; others fall back to internal AdamW. +Changes kept minimal relative to original Muon structure for integration simplicity. +""" + +import os + +import torch +import torch.distributed as dist + + +@torch.compile +def zeropower_via_newtonschulz5(G, steps=6, eps=1e-7): + assert len(G.shape) == 2 + a, b, c = (3.4445, -4.7750, 2.0315) + X = G.to(dtype=torch.bfloat16) + X /= X.norm() + eps + transposed = G.size(0) > G.size(1) + if transposed: + X = X.T + for _ in range(steps): + A = X @ X.T + B = b * A + c * A @ A + X = a * X + B @ X + if transposed: + X = X.T + return X + + +class AdaMuon(torch.optim.Optimizer): + def __init__( + self, + params, + lr=3e-4, + momentum=0.95, + ns_steps=6, + weight_decay=0.0, + adamw_lr=3e-4, + adamw_betas=(0.9, 0.95), + adamw_eps=1e-8, + adamw_wd=0.0, + eps=1e-8, + ): + defaults = dict( + lr=lr, + momentum=momentum, + ns_steps=ns_steps, + weight_decay=weight_decay, + adamw_lr=adamw_lr, + adamw_lr_ratio=adamw_lr / lr if lr != 0 else 1.0, + adamw_betas=adamw_betas, + adamw_eps=adamw_eps, + adamw_wd=adamw_wd, + eps=eps, + ) + param_list = list(params) + super().__init__(param_list, defaults) + + for p in param_list: + if p.ndim == 2 and p.size(0) < 10000: + self.state[p]["use_adamuon"] = True + else: + self.state[p]["use_adamuon"] = False + + if "WORLD_SIZE" in os.environ: + self.world_size = int(os.environ.get("WORLD_SIZE", 1)) + self.rank = int(os.environ.get("RANK", 0)) + else: + self.world_size = 1 + self.rank = 0 + + def step(self, closure=None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + for group in self.param_groups: + lr = group["lr"] + beta = group["momentum"] + ns_steps = group["ns_steps"] + weight_decay = group["weight_decay"] + eps = group["eps"] + + adamuon_params = [p for p in group["params"] if self.state[p]["use_adamuon"]] + + total_params = sum(p.numel() for p in adamuon_params) + updates_flat = torch.zeros(total_params, device=adamuon_params[0].device if adamuon_params else "cpu", dtype=torch.bfloat16) + curr_idx = 0 + + for i, p in enumerate(adamuon_params): + if i % self.world_size == self.rank: + g = p.grad + if g is None: + curr_idx += p.numel() + continue + g_view = g.view(g.size(0), -1) if g.ndim > 2 else g + state = self.state[p] + if "momentum_buffer" not in state: + state["momentum_buffer"] = torch.zeros_like(g_view) + if "rms_buffer" not in state: + state["rms_buffer"] = torch.zeros_like(g_view) + m_buf = state["momentum_buffer"] + v_buf = state["rms_buffer"] + m_buf.mul_(beta).add_(g_view) + sign_m = m_buf.sign() + O = zeropower_via_newtonschulz5(sign_m, steps=ns_steps) + O *= max(1, O.size(0) / O.size(1)) ** 0.5 + v_buf.mul_(beta).addcmul_(O, O, value=(1 - beta)) + O_hat = O / (v_buf.sqrt() + eps) + frob = O_hat.norm() + eps + gamma = 0.2 * (p.shape[0] * p.shape[1]) ** 0.5 / frob + update = (gamma * O_hat).to(dtype=torch.bfloat16) + updates_flat[curr_idx : curr_idx + p.numel()] = update.view(-1)[: p.numel()] + curr_idx += p.numel() + + if self.world_size > 1 and total_params > 0: + dist.all_reduce(updates_flat, op=dist.ReduceOp.SUM) + + curr_idx = 0 + for p in adamuon_params: + if p.grad is None: + curr_idx += p.numel() + continue + raw_update = updates_flat[curr_idx : curr_idx + p.numel()].view_as(p.data).to(p.data.dtype) + if weight_decay != 0: + raw_update = raw_update + weight_decay * p.data + p.data.add_(raw_update, alpha=-lr) + curr_idx += p.numel() + + # AdamW fallback for non-2D params + adamw_params = [p for p in group["params"] if not self.state[p]["use_adamuon"]] + if adamw_params: + aw_lr = group["adamw_lr_ratio"] * group["lr"] + beta1, beta2 = group["adamw_betas"] + aw_eps = group["adamw_eps"] + aw_wd = group["adamw_wd"] + for p in adamw_params: + g = p.grad + if g is None: + continue + state = self.state[p] + if "step" not in state: + state["step"] = 0 + state["moment1"] = torch.zeros_like(g) + state["moment2"] = torch.zeros_like(g) + state["step"] += 1 + step = state["step"] + m1 = state["moment1"] + m2 = state["moment2"] + m1.lerp_(g, 1 - beta1) + m2.lerp_(g.square(), 1 - beta2) + g_hat = m1 / (aw_eps + m2.sqrt()) + bc1 = 1 - beta1**step + bc2 = 1 - beta2**step + scale = bc1 / bc2**0.5 + if aw_wd != 0: + p.data.mul_(1 - aw_lr * aw_wd) + p.data.add_(g_hat, alpha=-aw_lr / scale) + return loss diff --git a/src/optimizers/main.py b/src/optimizers/main.py index ec35f1b..e951278 100644 --- a/src/optimizers/main.py +++ b/src/optimizers/main.py @@ -4,8 +4,8 @@ # from torch_optimizer import Shampoo sys.path.append("src/optimizers") import soap, muon -import taia import lora_rite +import adamuon, rmsspectral, adam_sania def get_optimizer(args, model): @@ -45,19 +45,48 @@ def get_optimizer(args, model): adamw_wd=args.weight_decay, momentum=args.momentum, ns_steps=args.ns_steps, + orth_algo=args.orth_algo, ) - elif args.optimizer == "taia": - optimizer = taia.TAIA( - taia_params=trainable_params, + elif args.optimizer == "adamuon": + optimizer = adamuon.AdaMuon( + params=trainable_params, + lr=args.lr, + momentum=args.momentum, + ns_steps=args.ns_steps, + weight_decay=args.weight_decay, + adamw_lr=args.adamw_lr, + adamw_betas=(args.beta1, args.beta2), + adamw_eps=args.eps, + adamw_wd=args.weight_decay, + eps=args.eps, + ) + elif args.optimizer == "rmsspectral": + optimizer = rmsspectral.RMSSpectral( + params=trainable_params, lr=args.lr, momentum=args.momentum, ns_steps=args.ns_steps, + weight_decay=args.weight_decay, adamw_lr=args.adamw_lr, adamw_betas=(args.beta1, args.beta2), adamw_eps=args.eps, adamw_wd=args.weight_decay, - lmo=args.lmo, - precondition_type=args.precondition_type, + eps=args.eps, + rms_power=getattr(args, "rms_power", 0.25), + ) + elif args.optimizer == "rmsspectral_sania": + optimizer = rmsspectral.RMSSpectral( + params=trainable_params, + lr=args.lr, + momentum=args.momentum, + ns_steps=args.ns_steps, + weight_decay=args.weight_decay, + adamw_lr=args.adamw_lr, + adamw_betas=(args.beta1, args.beta2), + adamw_eps=args.eps, + adamw_wd=args.weight_decay, + eps=args.eps, + rms_power=0.5, ) elif args.optimizer == "lora_rite": optimizer = lora_rite.LoRARite( @@ -67,6 +96,14 @@ def get_optimizer(args, model): lr=args.lr, weight_decay=args.weight_decay, ) + elif args.optimizer == "adam_sania": + optimizer = adam_sania.AdamSania( + params=trainable_params, + lr=args.lr, + betas=(args.beta1, args.beta2), + eps=args.eps, + weight_decay=args.weight_decay, + ) else: raise NotImplementedError(f"Wrong optimizer name {args.optimizer}") return optimizer diff --git a/src/optimizers/muon.py b/src/optimizers/muon.py index 977c64f..32d8b87 100644 --- a/src/optimizers/muon.py +++ b/src/optimizers/muon.py @@ -6,6 +6,7 @@ import os import torch +from itertools import repeat import torch.distributed as dist from typing import Callable @@ -35,6 +36,39 @@ def zeropower_via_newtonschulz5(G, steps=10, eps=1e-7): return X +# PolarExpress orthogonalization (polynomial approximation to the polar factor) +coeffs_list = [ + (8.28721201814563, -23.595886519098837, 17.300387312530933), + (4.107059111542203, -2.9478499167379106, 0.5448431082926601), + (3.9486908534822946, -2.908902115962949, 0.5518191394370137), + (3.3184196573706015, -2.488488024314874, 0.51004894012372), + (2.300652019954817, -1.6689039845747493, 0.4188073119525673), + (1.891301407787398, -1.2679958271945868, 0.37680408948524835), + (1.8750014808534479, -1.2500016453999487, 0.3750001645474248), + (1.875, -1.25, 0.375), # subsequent coeffs equal this numerically +] +# safety factor for numerical stability (but exclude last polynomial) +coeffs_list = [ + (a / 1.01, b / 1.01**3, c / 1.01**5) for (a, b, c) in coeffs_list[:-1] +] + [coeffs_list[-1]] + + +def polar_express(G: torch.Tensor, steps: int) -> torch.Tensor: + assert G.ndim >= 2 + X = G.bfloat16() + if G.size(-2) > G.size(-1): + X = X.mT # reduce FLOPs by operating on the smaller dimension first + X = X / (X.norm(dim=(-2, -1), keepdim=True) * 1.01 + 1e-7) + hs = coeffs_list[:steps] + list(repeat(coeffs_list[-1], steps - len(coeffs_list))) + for a, b, c in hs: + A = X @ X.mT + B = b * A + c * A @ A + X = a * X + B @ X # X <- aX + bX^3 + cX^5 + if G.size(-2) > G.size(-1): + X = X.mT + return X + + class Muon(torch.optim.Optimizer): """ Muon - MomentUm Orthogonalized by Newton-schulz @@ -69,6 +103,7 @@ def __init__( momentum=0.95, nesterov=True, ns_steps=6, + orth_algo: str = "ns", adamw_params=None, adamw_lr=3e-4, adamw_betas=(0.95, 0.95), @@ -80,6 +115,7 @@ def __init__( momentum=momentum, nesterov=nesterov, ns_steps=ns_steps, + orth_algo=orth_algo, adamw_lr=adamw_lr, adamw_lr_ratio=adamw_lr / lr, adamw_betas=adamw_betas, @@ -150,9 +186,12 @@ def step(self, closure: Callable = None): buf.mul_(momentum).add_(g) if group["nesterov"]: g = g.add(buf, alpha=momentum) - g = zeropower_via_newtonschulz5( - g, steps=group["ns_steps"], eps=group["adamw_eps"] - ) + if group.get("orth_algo", "ns") == "polar": + g = polar_express(g, steps=group["ns_steps"]) + else: + g = zeropower_via_newtonschulz5( + g, steps=group["ns_steps"], eps=group["adamw_eps"] + ) g *= max(1, g.size(0) / g.size(1)) ** 0.5 updates_flat[curr_idx : curr_idx + p.numel()] = g.flatten() curr_idx += p.numel() diff --git a/src/optimizers/rmsspectral.py b/src/optimizers/rmsspectral.py new file mode 100644 index 0000000..f642d49 --- /dev/null +++ b/src/optimizers/rmsspectral.py @@ -0,0 +1,180 @@ +"""RMSSpectral optimizer (AdaMuon variant with symmetric RMS spectral preconditioning). + +Generalization: use a power `p` for preconditioning instead of fixed 1/4. +For each 2D parameter: + M_t = beta * M_{t-1} + G_t + sign_m = Sign(M_t) + V_t = beta * V_{t-1} + (1 - beta) * (sign_m ⊙ sign_m) + R_t = V_t**p + eps (left preconditioning factor) + O_t = NS( sign_m / R_t ) (orthogonalization of left-preconditioned sign) + O_t *= spectral_scale (same normalization as AdaMuon) + O_hat = O_t / R_t (right preconditioning) + γ_t = 0.2 * sqrt(m*n) / ||O_hat||_F + W_{t+1} = W_t - lr * (γ_t * O_hat + weight_decay * W_t) + +Special cases: + p = 0.25 -> original RMSSpectral + p = 0.5 -> RMSSpectral-SANIA variant (stronger smoothing) +Non-2D params fall back to internal AdamW style update. +""" + +import os + +import torch +import torch.distributed as dist + + +@torch.compile +def zeropower_via_newtonschulz5(G, steps=6, eps=1e-7): + assert len(G.shape) == 2 + a, b, c = (3.4445, -4.7750, 2.0315) + X = G.to(dtype=torch.bfloat16) + X /= X.norm() + eps + transposed = G.size(0) > G.size(1) + if transposed: + X = X.T + for _ in range(steps): + A = X @ X.T + B = b * A + c * A @ A + X = a * X + B @ X + if transposed: + X = X.T + return X + + +class RMSSpectral(torch.optim.Optimizer): + def __init__( + self, + params, + lr=3e-4, + momentum=0.95, + ns_steps=6, + weight_decay=0.0, + adamw_lr=3e-4, + adamw_betas=(0.9, 0.95), + adamw_eps=1e-8, + adamw_wd=0.0, + eps=1e-8, + rms_power=0.25, + ): + defaults = dict( + lr=lr, + momentum=momentum, + ns_steps=ns_steps, + weight_decay=weight_decay, + adamw_lr=adamw_lr, + adamw_lr_ratio=adamw_lr / lr if lr != 0 else 1.0, + adamw_betas=adamw_betas, + adamw_eps=adamw_eps, + adamw_wd=adamw_wd, + eps=eps, + rms_power=rms_power, + ) + param_list = list(params) + super().__init__(param_list, defaults) + + for p in param_list: + if p.ndim == 2 and p.size(0) < 10000: + self.state[p]["use_adamuon"] = True + else: + self.state[p]["use_adamuon"] = False + + if "WORLD_SIZE" in os.environ: + self.world_size = int(os.environ.get("WORLD_SIZE", 1)) + self.rank = int(os.environ.get("RANK", 0)) + else: + self.world_size = 1 + self.rank = 0 + + def step(self, closure=None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + for group in self.param_groups: + lr = group["lr"] + beta = group["momentum"] + ns_steps = group["ns_steps"] + weight_decay = group["weight_decay"] + eps = group["eps"] + rms_power = group["rms_power"] + + adamuon_params = [p for p in group["params"] if self.state[p]["use_adamuon"]] + + total_params = sum(p.numel() for p in adamuon_params) + updates_flat = torch.zeros(total_params, device=adamuon_params[0].device if adamuon_params else "cpu", dtype=torch.bfloat16) + curr_idx = 0 + + for i, p in enumerate(adamuon_params): + if i % self.world_size == self.rank: + g = p.grad + if g is None: + curr_idx += p.numel() + continue + g_view = g.view(g.size(0), -1) if g.ndim > 2 else g + state = self.state[p] + if "momentum_buffer" not in state: + state["momentum_buffer"] = torch.zeros_like(g_view) + if "rms_buffer" not in state: + state["rms_buffer"] = torch.zeros_like(g_view) + m_buf = state["momentum_buffer"] + v_buf = state["rms_buffer"] + m_buf.mul_(beta).add_(g_view) + sign_m = m_buf.sign() + # Update second moment with sign variance + v_buf.mul_(beta).addcmul_(sign_m, sign_m, value=(1 - beta)) + rootp = v_buf.pow(rms_power) + eps # (V**p + eps) + pre_in = sign_m / rootp + O = zeropower_via_newtonschulz5(pre_in, steps=ns_steps) + O *= max(1, O.size(0) / O.size(1)) ** 0.5 + O_hat = O / rootp + frob = O_hat.norm() + eps + gamma = 0.2 * (p.shape[0] * p.shape[1]) ** 0.5 / frob + update = (gamma * O_hat).to(dtype=torch.bfloat16) + updates_flat[curr_idx : curr_idx + p.numel()] = update.view(-1)[: p.numel()] + curr_idx += p.numel() + + if self.world_size > 1 and total_params > 0: + dist.all_reduce(updates_flat, op=dist.ReduceOp.SUM) + + curr_idx = 0 + for p in adamuon_params: + if p.grad is None: + curr_idx += p.numel() + continue + raw_update = updates_flat[curr_idx : curr_idx + p.numel()].view_as(p.data).to(p.data.dtype) + if weight_decay != 0: + raw_update = raw_update + weight_decay * p.data + p.data.add_(raw_update, alpha=-lr) + curr_idx += p.numel() + + # AdamW fallback for non-2D params + adamw_params = [p for p in group["params"] if not self.state[p]["use_adamuon"]] + if adamw_params: + aw_lr = group["adamw_lr_ratio"] * group["lr"] + beta1, beta2 = group["adamw_betas"] + aw_eps = group["adamw_eps"] + aw_wd = group["adamw_wd"] + for p in adamw_params: + g = p.grad + if g is None: + continue + state = self.state[p] + if "step" not in state: + state["step"] = 0 + state["moment1"] = torch.zeros_like(g) + state["moment2"] = torch.zeros_like(g) + state["step"] += 1 + step = state["step"] + m1 = state["moment1"] + m2 = state["moment2"] + m1.lerp_(g, 1 - beta1) + m2.lerp_(g.square(), 1 - beta2) + g_hat = m1 / (aw_eps + m2.sqrt()) + bc1 = 1 - beta1**step + bc2 = 1 - beta2**step + scale = bc1 / bc2**0.5 + if aw_wd != 0: + p.data.mul_(1 - aw_lr * aw_wd) + p.data.add_(g_hat, alpha=-aw_lr / scale) + return loss diff --git a/src/optimizers/taia.py b/src/optimizers/taia.py deleted file mode 100644 index e83705f..0000000 --- a/src/optimizers/taia.py +++ /dev/null @@ -1,348 +0,0 @@ -""" -Here is an original implementation of Muon. -Source: https://github.com/KellerJordan/modded-nanogpt -""" - -import os - -from numpy import dtype -import torch -import torch.distributed as dist -from typing import Callable - - -def zeropower_via_newtonschulz5(G, steps=10, eps=1e-7): - """ - Newton-Schulz iteration to compute the zeroth power / orthogonalization of G. We opt to use a - quintic iteration whose coefficients are selected to maximize the slope at zero. For the purpose - of minimizing steps, it turns out to be empirically effective to keep increasing the slope at - zero even beyond the point where the iteration no longer converges all the way to one everywhere - on the interval. This iteration therefore does not produce UV^T but rather something like US'V^T - where S' is diagonal with S_{ii}' \sim Uniform(0.5, 1.5), which turns out not to hurt model - performance at all relative to UV^T, where USV^T = G is the SVD. - """ - assert len(G.shape) == 2 - a, b, c = (3.4445, -4.7750, 2.0315) - X = G.bfloat16() - X /= X.norm() + eps # ensure top singular value <= 1 - if G.size(0) > G.size(1): - X = X.T - for _ in range(steps): - A = X @ X.T - B = b * A + c * A @ A - X = a * X + B @ X - if G.size(0) > G.size(1): - X = X.T - - return X.to(dtype=G.dtype, device=G.device) - - -class TAIA(torch.optim.Optimizer): - """ - Arguments: - muon_params: The parameters to be optimized by Muon. - lr: The learning rate. The updates will have spectral norm of `lr`. (0.02 is a good default) - momentum: The momentum used by the internal SGD. (0.95 is a good default) - nesterov: Whether to use Nesterov-style momentum in the internal SGD. (recommended) - ns_steps: The number of Newton-Schulz iterations to run. (6 is probably always enough) - adamw_params: The parameters to be optimized by AdamW. Any parameters in `muon_params` which are - {0, 1}-D or are detected as being the embed or lm_head will be optimized by AdamW as well. - adamw_lr: The learning rate for the internal AdamW. - adamw_betas: The betas for the internal AdamW. - adamw_eps: The epsilon for the internal AdamW. - adamw_wd: The weight decay for the internal AdamW. - """ - - def __init__( - self, - taia_params, - lr=0.02, - momentum=0.95, - nesterov=True, - ns_steps=6, - adamw_params=None, - adamw_lr=3e-4, - adamw_betas=(0.95, 0.95), - adamw_eps=1e-8, - adamw_wd=0, - lmo=None, - precondition_type="norm", - ): - defaults = dict( - lr=lr, - momentum=momentum, - nesterov=nesterov, - ns_steps=ns_steps, - adamw_lr=adamw_lr, - adamw_lr_ratio=adamw_lr / lr, - adamw_betas=adamw_betas, - adamw_eps=adamw_eps, - adamw_wd=adamw_wd, - lmo=lmo, - precondition_type=precondition_type, - ) - - params = list(taia_params) - adamw_params = list(adamw_params) if adamw_params is not None else [] - params.extend(adamw_params) - super().__init__(params, defaults) - - # Sort parameters into those for which we will use Muon, and those for which we will not - for p in taia_params: - # Use Muon for every parameter in muon_params which is >= 2D and doesn't look like an embedding or head layer - if p.ndim >= 2 and p.size(0) < 10000: - self.state[p]["use_taia"] = True - else: - self.state[p]["use_taia"] = False - for p in adamw_params: - # Do not use Muon for parameters in adamw_params - self.state[p]["use_taia"] = False - - if "WORLD_SIZE" in os.environ: - self.world_size = int(os.environ["WORLD_SIZE"]) - self.rank = int(os.environ["RANK"]) - else: - self.world_size = 1 - self.rank = 0 - - def step(self, closure: Callable = None): - """ - Performs a single optimization step. - - Arguments: - closure (`Callable`, *optional*): A closure that reevaluates the model and returns the loss. - """ - loss = None - if closure is not None: - loss = closure() - for group in self.param_groups: - ############################ - # TAIA # - ############################ - - params = [p for p in group["params"] if self.state[p]["use_taia"]] - lr = group["lr"] - momentum = group["momentum"] - - # generate weight updates in distributed fashion - total_params = sum(p.numel() for p in params) - updates_flat = torch.zeros( - total_params, device="cuda", dtype=torch.bfloat16 - ) - curr_idx = 0 - for i, p in enumerate(params): - # luckily this will perfectly distribute a transformer with multiple of 4 layers to 8 GPUs - if i % self.world_size == self.rank: - g = p.grad - if g.ndim > 2: - g = g.view(g.size(0), -1) - assert g is not None - state = self.state[p] - if "step " not in state: - state["step"] = 0 - if "momentum_buffer" not in state: - state["momentum_buffer"] = torch.zeros_like(g) - if "exp_avg" not in state: - state["exp_avg"] = torch.zeros_like(g) - if "prec_L" not in state and group["precondition_type"] == "fisher": - # state["prec_L"] = shampoo_eps * torch.eye(g.size(0), device=g.device) - # state["prec_R"] = shampoo_eps * torch.eye(g.size(1), device=g.device) - state["prec_L"] = torch.zeros( - g.size(0), g.size(0), device=g.device - ) - state["prec_R"] = torch.zeros( - g.size(1), g.size(1), device=g.device - ) - if ( - "exp_avg_sq" not in state - and group["precondition_type"] == "adam" - ): - state["exp_avg_sq"] = torch.zeros_like(g) - - # First momentum with bias correction - exp_avg = state["exp_avg"] - exp_avg.mul_(momentum).add_(g, alpha=1.0 - momentum) - - state["step"] += 1 - bias_correction1 = 1.0 - momentum ** (state["step"]) - exp_avg_corrected = exp_avg / bias_correction1 - if group["precondition_type"] == "fisher": - H_L = state["prec_L"].clone() - H_R = state["prec_R"].clone() - # print(H_R) - state["prec_L"].add_(g @ g.T) - state["prec_R"].add_(g.T @ g) - elif group["precondition_type"] == "adam": - beta2 = group["adamw_betas"][1] - state["exp_avg_sq"].mul_(beta2).addcmul_( - g, g, value=1.0 - beta2 - ) - bias_correction2 = 1.0 - beta2 ** state["step"] - D = ( - state["exp_avg_sq"] - .sqrt() - .add_(group["adamw_eps"]) - .div_(bias_correction2**0.5) - ) - exp_avg_corrected /= D.sqrt() - elif group["precondition_type"] == "adam_sania": - beta2 = group["adamw_betas"][1] - state["exp_avg_sq"].mul_(beta2).addcmul_( - g, g, value=1.0 - beta2 - ) - bias_correction2 = 1.0 - beta2 ** state["step"] - D = ( - state["exp_avg_sq"] - .add_(group["adamw_eps"]) - .div_(bias_correction2) - ) # remove square root - exp_avg_corrected /= D.sqrt() - - if group["precondition_type"] == "norm": - L = g.norm(dim=0, keepdim=True) - L = torch.where(L == 0, 1e-8, L) - exp_avg_corrected /= L - elif group["precondition_type"] == "fisher": - # L = torch.linalg.cholesky(H_L, upper=True) - # R = torch.linalg.cholesky(H_R, upper=False) - # L_inv = torch.linalg.inv(L).to(dtype=g.dtype, device=g.device) - # R_inv = torch.linalg.inv(R).to(dtype=g.dtype, device=g.device) - U_L, sigma_L, _ = torch.linalg.svd(H_L) - U_R, sigma_R, _ = torch.linalg.svd(H_R) - for i in range(sigma_L.size(0)): - sigma_L[i] = 1.0 / sigma_L[i] if sigma_L[i] > 1e-10 else 0 - for i in range(sigma_R.size(0)): - sigma_R[i] = 1.0 / sigma_R[i] if sigma_R[i] > 1e-10 else 0 - L_inv = U_L @ torch.diag(sigma_L**1 / 8) - R_inv = torch.diag(sigma_R**1 / 8) @ U_R.T - exp_avg_corrected = L_inv.T @ exp_avg_corrected @ R_inv.T - - if group["lmo"] == "spectral": - exp_avg_corrected = zeropower_via_newtonschulz5( - exp_avg_corrected, - steps=group["ns_steps"], - eps=group["adamw_eps"], - ) - exp_avg_corrected *= max(1, g.size(0) / g.size(1)) ** 0.5 - - if group["precondition_type"] == "norm": - exp_avg_corrected /= L - elif group["precondition_type"] == "fisher": - exp_avg_corrected = L_inv @ exp_avg_corrected @ R_inv - elif group["precondition_type"] in ["adam", "adam_sania"]: - exp_avg_corrected /= D.sqrt() - # g /= D.sqrt() - - updates_flat[curr_idx : curr_idx + p.numel()] = ( - exp_avg_corrected.flatten() - ) - curr_idx += p.numel() - - # sync updates across devices. we are not memory-constrained so can do this simple deserialization - if self.world_size > 1: - dist.all_reduce(updates_flat, op=dist.ReduceOp.SUM) - - # deserialize and apply updates - curr_idx = 0 - for p in params: - exp_avg_corrected = ( - updates_flat[curr_idx : curr_idx + p.numel()] - .view_as(p.data) - .type_as(p.data) - ) - - p.data.add_(exp_avg_corrected, alpha=-lr) - - curr_idx += p.numel() - - ############################ - # AdamW backup # - ############################ - - params = [p for p in group["params"] if not self.state[p]["use_taia"]] - lr = ( - group["adamw_lr_ratio"] * group["lr"] - ) # in order for lr schedule to work - beta1, beta2 = group["adamw_betas"] - eps = group["adamw_eps"] - weight_decay = group["adamw_wd"] - - for p in params: - g = p.grad - assert g is not None - state = self.state[p] - if "step" not in state: - state["step"] = 0 - state["moment1"] = torch.zeros_like(g) - state["moment2"] = torch.zeros_like(g) - state["step"] += 1 - step = state["step"] - buf1 = state["moment1"] - buf2 = state["moment2"] - buf1.lerp_(g, 1 - beta1) - buf2.lerp_(g.square(), 1 - beta2) - - g = buf1 / (eps + buf2.sqrt()) - - bias_correction1 = 1 - beta1**step - bias_correction2 = 1 - beta2**step - scale = bias_correction1 / bias_correction2**0.5 - p.data.mul_(1 - lr * weight_decay) - p.data.add_(g, alpha=-lr / scale) - - return loss - - -def separate_params(param_groups): - param_groups_2d = [] - param_groups_non2d = [] - total_param_2d_count = 0 - total_param_non2d_count = 0 - - # Check if param_groups is a list of dicts or list of params - if ( - isinstance(param_groups, list) and isinstance(param_groups[0], dict) - ) or isinstance(param_groups, dict): - if isinstance(param_groups, dict): - param_groups = [param_groups] - # param_groups is a list of dicts - for group in param_groups: - ( - params_2d, - params_non2d, - param_2d_count, - param_non2d_count, - ) = separate_params(group["params"]) - param_group_2d = {"params": params_2d} - param_group_non2d = {"params": params_non2d} - # Copy the group dict and replace the 'params' key with the separated params - for k in group.keys(): - if k != "params": - param_group_2d[k] = group[k] - param_group_non2d[k] = group[k] - - param_groups_2d.append(param_group_2d) - param_groups_non2d.append(param_group_non2d) - total_param_2d_count += param_2d_count - total_param_non2d_count += param_non2d_count - - return ( - param_groups_2d, - param_groups_non2d, - total_param_2d_count, - total_param_non2d_count, - ) - - elif isinstance(param_groups, list) and isinstance(param_groups[0], torch.Tensor): - params_2d = [] - params_non2d = [] - param_group = param_groups - # param_group is a list of param tensors - for param in param_group: - if param.ndim == 2: - params_2d.append(param) - else: - params_non2d.append(param) - return params_2d, params_non2d, len(params_2d), len(params_non2d) - else: - breakpoint() \ No newline at end of file diff --git a/src/utils.py b/src/utils.py index f06a554..51db1b8 100644 --- a/src/utils.py +++ b/src/utils.py @@ -92,7 +92,7 @@ def get_run_name(args, parser, tuning=False): "weight_init", "ns_steps", ] - if args.optimizer in ["taia", "adam-sania"]: + if args.optimizer in ["adam-sania"]: ignore_args_tuning.append("scale") # Get the default values defaults = vars(parser.parse_args([])) diff --git a/wandb_results/glue_results.csv b/wandb_results/glue_results.csv new file mode 100644 index 0000000..ee0e612 --- /dev/null +++ b/wandb_results/glue_results.csv @@ -0,0 +1,91 @@ +ft_strategy,dataset,optimizer,best_value,lr,run_id +LoRA,cola,adamuon,0.5174,0.0003,4b7nfsyh +LoRA,cola,adamw,0.4759,5e-05,iwxks9su +LoRA,cola,muon,0.4992,0.0002,byd47ni5 +LoRA,cola,rmsspectral,0.5327,0.001,nlvqnx3e +LoRA,cola,rmsspectral_sania,0.4990,0.0003,i7kl2ixl +LoRA,mnli,adamuon,0.7836,0.001,z7acsdpy +LoRA,mnli,adamw,0.7352,0.0002,9psqck5s +LoRA,mnli,muon,0.7237,0.0002,fg9a6v70 +LoRA,mnli,rmsspectral,0.7776,0.001,a1w9upbs +LoRA,mnli,rmsspectral_sania,0.7780,0.001,znw4ykf4 +LoRA,mrpc,adamuon,0.8480,0.001,fpxumwci +LoRA,mrpc,adamw,0.8309,0.0001,p9keujfk +LoRA,mrpc,muon,0.8578,3e-05,dp194vz7 +LoRA,mrpc,rmsspectral,0.8505,0.0003,ae4m02ln +LoRA,mrpc,rmsspectral_sania,0.8480,0.0001,qjhd3gkz +LoRA,qnli,adamuon,0.8739,0.001,uwtwhbk4 +LoRA,qnli,adamw,0.8567,0.0002,pvglqkhd +LoRA,qnli,muon,0.8501,0.0002,qat8x1zz +LoRA,qnli,rmsspectral,0.8682,0.001,z06nmiax +LoRA,qnli,rmsspectral_sania,0.8781,0.001,fwu3ncxw +LoRA,qqp,adamuon,0.8584,0.001,yz0po716 +LoRA,qqp,adamw,0.8498,0.0002,ux66i5vg +LoRA,qqp,muon,0.8322,0.0002,gcuquvh1 +LoRA,qqp,rmsspectral,0.8585,0.001,qdhjxozy +LoRA,qqp,rmsspectral_sania,0.8571,0.001,qez9vuq9 +LoRA,rte,adamuon,0.6679,0.001,c85hqz1y +LoRA,rte,adamw,0.6390,0.0001,zyqmv4my +LoRA,rte,muon,0.6606,3e-05,xne5jqhy +LoRA,rte,rmsspectral,0.6534,0.0001,kl4l58bu +LoRA,rte,rmsspectral_sania,0.6209,0.0003,t29o81n2 +LoRA,sst2,adamuon,0.9037,0.0003,sin6eidh +LoRA,sst2,adamw,0.8956,5e-05,42sxr048 +LoRA,sst2,muon,0.8865,0.0001,ojk5qpwl +LoRA,sst2,rmsspectral,0.8968,0.001,6gdel29s +LoRA,sst2,rmsspectral_sania,0.9060,0.001,r15xs3jh +LoRA,stsb,adamuon,0.8427,0.0003,7q67xmid +LoRA,stsb,adamw,0.8473,0.0001,nbmlsz2d +LoRA,stsb,muon,0.8392,5e-05,65wj4b6q +LoRA,stsb,rmsspectral,0.8344,0.0003,3ls2shjc +LoRA,stsb,rmsspectral_sania,0.8422,0.0003,tvmyhqhe +LoRA,wnli,adamuon,n/a,, +LoRA,wnli,adamw,0.1549,0.0002,i6u3eni4 +LoRA,wnli,muon,0.1408,3e-05,usz676dn +LoRA,wnli,rmsspectral,n/a,, +LoRA,wnli,rmsspectral_sania,n/a,, +Full,cola,adamuon,0.4770,0.0003,y28kf6jj +Full,cola,adamw,0.5880,5e-05,mlms7l8a +Full,cola,muon,0.5753,1e-05,eyj8ptfw +Full,cola,rmsspectral,0.4800,0.0003,zhihfmmm +Full,cola,rmsspectral_sania,0.4587,0.0003,5txm0lij +Full,mnli,adamuon,0.8109,0.001,9uxsh67v +Full,mnli,adamw,0.8411,5e-05,yq7zert3 +Full,mnli,muon,0.8068,0.0001,dqlmczd0 +Full,mnli,rmsspectral,0.8116,0.001,9va1z6z9 +Full,mnli,rmsspectral_sania,0.8122,0.001,tdemrqpq +Full,mrpc,adamuon,0.8260,2e-05,8tu5nyso +Full,mrpc,adamw,0.8578,5e-05,w73224gp +Full,mrpc,muon,0.8456,5e-05,frr1xoll +Full,mrpc,rmsspectral,0.7819,2e-05,h3g02gb2 +Full,mrpc,rmsspectral_sania,0.8137,2e-05,x38x40gj +Full,qnli,adamuon,0.8757,0.001,ammivuag +Full,qnli,adamw,0.9189,5e-05,5w69qa3t +Full,qnli,muon,0.9055,0.0001,dzj948tb +Full,qnli,rmsspectral,0.8794,0.001,ytknk9ot +Full,qnli,rmsspectral_sania,0.8788,0.001,mt0zgqp0 +Full,qqp,adamuon,0.8968,0.001,wovf4ndr +Full,qqp,adamw,0.9063,5e-05,i42jbtu2 +Full,qqp,muon,0.8958,0.0002,h10vr5m6 +Full,qqp,rmsspectral,0.8985,0.001,5oplujg6 +Full,qqp,rmsspectral_sania,0.8978,0.001,mdcko8wq +Full,rte,adamuon,0.5740,0.001,ex4f642q +Full,rte,adamw,0.6823,5e-05,qxy7cvns +Full,rte,muon,0.6895,1e-05,mnnb4apy +Full,rte,rmsspectral,0.6137,2e-05,s7so3pwi +Full,rte,rmsspectral_sania,0.5740,0.0003,illoegwy +Full,sst2,adamuon,0.8968,0.001,1hz8cgdd +Full,sst2,adamw,0.9289,1e-05,6rgwgz7k +Full,sst2,muon,0.9197,0.0001,5fqp432q +Full,sst2,rmsspectral,0.9002,0.0003,njsjo9y2 +Full,sst2,rmsspectral_sania,0.9071,0.001,tvaz8ja6 +Full,stsb,adamuon,0.8300,0.0003,ap215vjt +Full,stsb,adamw,0.8932,5e-05,wk3otfqr +Full,stsb,muon,0.8903,0.0001,241ef2y7 +Full,stsb,rmsspectral,0.8137,0.0003,syzfj8ra +Full,stsb,rmsspectral_sania,0.8142,0.0003,v9qe6dhh +Full,wnli,adamuon,n/a,, +Full,wnli,adamw,0.0986,0.0002,irj3dupi +Full,wnli,muon,0.1408,3e-05,raho3g5b +Full,wnli,rmsspectral,n/a,, +Full,wnli,rmsspectral_sania,n/a,, diff --git a/wandb_results/glue_table.tex b/wandb_results/glue_table.tex new file mode 100644 index 0000000..f3c43bd --- /dev/null +++ b/wandb_results/glue_table.tex @@ -0,0 +1,26 @@ +\begin{table}[h!] +\centering +\scriptsize +\setlength{\tabcolsep}{3pt} +\renewcommand{\arraystretch}{1.1} +\captionof{table}{GLUE: datasets are columns with the corresponding metric; \texttt{ALL} is the average over tasks. Best per task in bold.} +\resizebox{0.95\linewidth}{!}{ +\begin{tabular}{l lcccccccc|c} +& & \begin{tabular}{@{}c@{}}CoLA\\Matthews\end{tabular} & \begin{tabular}{@{}c@{}}MNLI\\Acc\end{tabular} & \begin{tabular}{@{}c@{}}MRPC\\Acc\end{tabular} & \begin{tabular}{@{}c@{}}QNLI\\Acc\end{tabular} & \begin{tabular}{@{}c@{}}QQP\\Acc\end{tabular} & \begin{tabular}{@{}c@{}}RTE\\Acc\end{tabular} & \begin{tabular}{@{}c@{}}SST-2\\Acc\end{tabular} & \begin{tabular}{@{}c@{}}STS-B\\Comb.\end{tabular} & \begin{tabular}{@{}c@{}}ALL\\Avg\end{tabular} \\ +\midrule +\multirow{5}{*}{LoRA} & \texttt{AdaMuon} & 0.5174 & \textbf{0.7836} & 0.8480 & 0.8739 & 0.8584 & \textbf{0.6679} & 0.9037 & 0.8427 & \textbf{0.7869} \\ + & \texttt{AdamW} & 0.4759 & 0.7352 & 0.8309 & 0.8567 & 0.8498 & 0.6390 & 0.8956 & \textbf{0.8473} & 0.7663 \\ + & \texttt{Muon} & 0.4992 & 0.7237 & \textbf{0.8578} & 0.8501 & 0.8322 & 0.6606 & 0.8865 & 0.8392 & 0.7687 \\ + & \texttt{RMSSpectral} & \textbf{0.5327} & 0.7776 & 0.8505 & 0.8682 & \textbf{0.8585} & 0.6534 & 0.8968 & 0.8344 & 0.7840 \\ + & \texttt{RMSSpectral-SANIA} & 0.4990 & 0.7780 & 0.8480 & \textbf{0.8781} & 0.8571 & 0.6209 & \textbf{0.9060} & 0.8422 & 0.7787 \\ +\midrule +\multirow{5}{*}{Full} & \texttt{AdaMuon} & 0.4770 & 0.8109 & 0.8260 & 0.8757 & 0.8968 & 0.5740 & 0.8968 & 0.8300 & 0.7734 \\ + & \texttt{AdamW} & \textbf{0.5880} & \textbf{0.8411} & \textbf{0.8578} & \textbf{0.9189} & \textbf{0.9063} & 0.6823 & \textbf{0.9289} & \textbf{0.8932} & \textbf{0.8271} \\ + & \texttt{Muon} & 0.5753 & 0.8068 & 0.8456 & 0.9055 & 0.8958 & \textbf{0.6895} & 0.9197 & 0.8903 & 0.8161 \\ + & \texttt{RMSSpectral} & 0.4800 & 0.8116 & 0.7819 & 0.8794 & 0.8985 & 0.6137 & 0.9002 & 0.8137 & 0.7724 \\ + & \texttt{RMSSpectral-SANIA} & 0.4587 & 0.8122 & 0.8137 & 0.8788 & 0.8978 & 0.5740 & 0.9071 & 0.8142 & 0.7696 \\ +\bottomrule +\end{tabular} +} +\label{tab:glue_results_transposed} +\end{table} \ No newline at end of file