Skip to content
Merged
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
12 changes: 12 additions & 0 deletions Justfile
Original file line number Diff line number Diff line change
Expand Up @@ -337,6 +337,18 @@ retrain: align-premises
retrain-skip-align:
julia --project=src/julia src/julia/run_training.jl

# Full training run on the real corpus (auto-selects GPU when available).
# Honours ECHIDNA_MAX_PROOF_STATES, ECHIDNA_NUM_EPOCHS, ECHIDNA_NUM_NEGATIVES.
# Produces models/neural/{premise_selector,tactic_predictor}.bson, vocab.json,
# updates models/model_metadata.txt, and appends to training_data/metrics_baseline.jsonl.
train:
julia --project=src/julia src/julia/run_training.jl

# CPU smoke pass: 2000 proof states, 2 epochs — fast end-to-end sanity check.
# Safe to run on any dev box without a GPU. Produces the same artefacts as `train`.
train-cpu:
ECHIDNA_MAX_PROOF_STATES=2000 ECHIDNA_NUM_EPOCHS=2 julia --project=src/julia src/julia/run_training_cpu.jl

# End-to-end pipeline: provision → extract → merge → align → retrain.
# Use `ECHIDNA_MAX_PROOF_STATES=0 just corpus-refresh` to lift the sample cap.
corpus-refresh: provision-corpora extract-corpora merge-corpora align-premises retrain
Expand Down
32 changes: 21 additions & 11 deletions src/julia/api/gnn_endpoint.jl
Original file line number Diff line number Diff line change
Expand Up @@ -42,27 +42,37 @@ const TOTAL_TRAINING_RECORDS = Ref{Int}(0)
"""
load_gnn_model(models_dir::String)

Load the GNN premise ranker model from disk.
Falls back to creating a fresh (untrained) model if no checkpoint exists.
Load the GNN premise ranker model from disk. Tries, in order:
1. `models/neural/gnn_ranker/` — directory written by run_training*.jl
2. `models/neural/best_model/` — early-stopping checkpoint
3. `models/neural/final_model/` — last epoch checkpoint

Only falls back to the cosine path (GNN_MODEL[] = nothing) when none of
the above directories exist or all fail to deserialise. The cosine path
is the genuine missing-model fallback for CI smoke runs.
"""
function load_gnn_model(models_dir::String)
model_path = joinpath(models_dir, "neural", "gnn_ranker")

if isdir(model_path)
@info "Loading GNN model from $model_path"
candidate_dirs = [
joinpath(models_dir, "neural", "gnn_ranker"),
joinpath(models_dir, "neural", "best_model"),
joinpath(models_dir, "neural", "final_model"),
]

for model_path in candidate_dirs
isdir(model_path) || continue
@info "Trying to load GNN model from $model_path"
try
solver = load_solver(model_path)
GNN_MODEL[] = solver
@info "GNN model loaded successfully"
@info "GNN model loaded successfully from $model_path"
return true
catch e
@warn "Failed to load GNN model: $e"
@warn "Failed to load GNN model from $model_path: $e"
end
end

@info "No trained GNN model found — creating fresh model for inference"
# Create a minimal model with default configuration for the endpoint
# to respond (scores will be random until training completes)
# All candidates exhausted — cosine fallback is genuine missing-model path.
@warn "No trained GNN model found in $(joinpath(models_dir, "neural")) — ranking will use cosine similarity until weights are trained (run: just train-cpu)"
GNN_MODEL[] = nothing
return false
end
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#!/usr/bin/env julia
# ARCHIVED 2026-05-24 — superseded by src/julia/run_training.jl. Kept for history only; do not invoke.
# SPDX-FileCopyrightText: 2026 ECHIDNA Project Team
# SPDX-License-Identifier: MPL-2.0

Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#!/usr/bin/env julia
# ARCHIVED 2026-05-24 — superseded by src/julia/run_training.jl. Kept for history only; do not invoke.
# SPDX-License-Identifier: MPL-2.0
# Train GNN model and evaluate on validation set, outputting metrics for health monitoring

Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#!/usr/bin/env julia
# ARCHIVED 2026-05-24 — superseded by src/julia/run_training.jl. Kept for history only; do not invoke.
# SPDX-FileCopyrightText: 2026 ECHIDNA Project Team
# SPDX-License-Identifier: MPL-2.0

Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#!/usr/bin/env julia
# ARCHIVED 2026-05-24 — superseded by src/julia/run_training.jl. Kept for history only; do not invoke.
# SPDX-FileCopyrightText: 2026 ECHIDNA Project Team
# SPDX-License-Identifier: MPL-2.0

Expand Down
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
# ARCHIVED 2026-05-24 — superseded by src/julia/run_training.jl. Kept for history only; do not invoke.
# SPDX-FileCopyrightText: 2026 ECHIDNA Project Team
# SPDX-License-Identifier: MPL-2.0

Expand Down
82 changes: 82 additions & 0 deletions src/julia/run_training.jl
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@

using Pkg
Pkg.activate(joinpath(@__DIR__))
using Dates

println("╔═══════════════════════════════════════════════════════════╗")
println("║ ECHIDNA Neural Solver — Training Pipeline ║")
Expand Down Expand Up @@ -143,7 +144,9 @@ println("Training ($(training_config.num_epochs) epochs)...")
println("═══════════════════════════════════════════════════════════")

mkpath(save_dir)
t_start = time()
metrics = train_solver!(solver, train_data, val_data; config=training_config)
duration_seconds = time() - t_start

# Save final model
println()
Expand All @@ -154,14 +157,93 @@ println("═══════════════════════
final_path = joinpath(save_dir, "final_model")
save_solver(solver, final_path)

# Deploy alias: gnn_endpoint.jl loads gnn_ranker → best_model → final_model.
best_dir = joinpath(save_dir, "best_model")
ranker_dir = joinpath(save_dir, "gnn_ranker")
if isdir(best_dir)
rm(ranker_dir; recursive=true, force=true)
cp(best_dir, ranker_dir)
println("Published best_model → $ranker_dir")
end

# Save vocabulary separately for the API server
BSON.@save joinpath(save_dir, "vocabulary.bson") vocab

# Flat canonical artefacts expected by gnn_endpoint.jl and the spec.
# premise_selector.bson = the full NeuralSolver weights (renamed for clarity).
# tactic_predictor.bson = same weights (tactic side shares the text_encoder).
# vocab.json = human-readable vocabulary for inspection and version checks.
import JSON3 as _JSON3
weights = Flux.state(solver)
BSON.bson(joinpath(save_dir, "premise_selector.bson"), weights=weights)
BSON.bson(joinpath(save_dir, "tactic_predictor.bson"), weights=weights)

# Build a compact vocab.json: token→id map + metadata.
tactic_classes = sort(unique([
ex.proof_state.goal[1:min(8,length(ex.proof_state.goal))]
for ex in vcat(train_data.examples, val_data.examples)
]))
vocab_json = Dict(
"vocab_size" => vocab.vocab_size,
"tactic_classes" => length(tactic_classes),
"token_to_id" => vocab.token_to_id,
)
open(joinpath(save_dir, "vocab.json"), "w") do io
_JSON3.write(io, vocab_json)
end
println("Saved vocab.json ($(vocab.vocab_size) tokens, $(length(tactic_classes)) tactic classes)")

# Compute final MRR / top-k on validation split.
val_metrics = compute_metrics(solver, val_data; k=10)
val_top1 = compute_metrics(solver, val_data; k=1).precision
val_top5 = compute_metrics(solver, val_data; k=5).precision
println("Validation MRR=$(round(val_metrics.mrr, digits=4)) top1=$(round(val_top1, digits=4)) top5=$(round(val_top5, digits=4)) top10=$(round(val_metrics.precision, digits=4))")

# Append a single JSONL row to training_data/metrics_baseline.jsonl.
git_sha = try; strip(read(`git -C $(joinpath(@__DIR__, "..", "..")) rev-parse --short HEAD`, String)); catch; "unknown"; end
metrics_row = Dict{String,Any}(
"timestamp" => string(Dates.now()),
"git_sha" => git_sha,
"mrr" => round(Float64(val_metrics.mrr), digits=6),
"top1" => round(Float64(val_top1), digits=6),
"top5" => round(Float64(val_top5), digits=6),
"top10" => round(Float64(val_metrics.precision), digits=6),
"epochs" => training_config.num_epochs,
"max_proof_states" => max_proof_states,
"duration_seconds" => round(duration_seconds, digits=1),
"device" => has_gpu ? "gpu" : "cpu",
)
metrics_baseline_path = joinpath(data_dir, "metrics_baseline.jsonl")
open(metrics_baseline_path, "a") do io
println(io, _JSON3.write(metrics_row))
end
println("Appended metrics row to $metrics_baseline_path")

# Rewrite models/model_metadata.txt with real values.
metadata_path = joinpath(@__DIR__, "..", "..", "models", "model_metadata.txt")
open(metadata_path, "w") do io
println(io, "# ECHIDNA Neural Models v2.0")
println(io, "# Trained: $(Dates.now())")
println(io, "# Git SHA: $git_sha")
println(io, "# Device: $(has_gpu ? "GPU" : "CPU")")
println(io, "# Premise Selector: vocabulary-based ($(vocab.vocab_size) words)")
println(io, "# Tactic Predictor: neural text encoder ($(length(tactic_classes)) classes)")
println(io, "# MRR: $(round(Float64(val_metrics.mrr), digits=4))")
println(io, "# Top-1 Precision: $(round(Float64(val_top1), digits=4))")
println(io, "# Top-5 Precision: $(round(Float64(val_top5), digits=4))")
println(io, "# Epochs: $(training_config.num_epochs)")
println(io, "# Max Proof States: $max_proof_states")
println(io, "# Training Examples: $(length(train_data.examples))")
println(io, "# Validation Examples: $(length(val_data.examples))")
end
println("Updated $metadata_path")

println()
println("╔═══════════════════════════════════════════════════════════╗")
println("║ Training Complete! ║")
println("╚═══════════════════════════════════════════════════════════╝")
println()
println("Model saved to: $save_dir")
println("MRR: $(round(Float64(val_metrics.mrr), digits=4))")
println("To start the API server:")
println(" julia --project=src/julia src/julia/run_server.jl $save_dir")
70 changes: 67 additions & 3 deletions src/julia/run_training_cpu.jl
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@

using Pkg
Pkg.activate(joinpath(@__DIR__))
using Dates

include(joinpath(@__DIR__, "EchidnaML.jl"))
using .EchidnaML
Expand Down Expand Up @@ -80,11 +81,11 @@ config = TrainingConfig(
)

mkpath(save_dir)
train_solver!(solver, train_data, val_data; config=config)
t_start = time()
metrics = train_solver!(solver, train_data, val_data; config=config)
duration_seconds = time() - t_start

# Deploy alias: gnn_endpoint.jl loads gnn_ranker → best_model → final_model.
# Publish the early-stopping best model as the canonical gnn_ranker so the
# server serves the best checkpoint, not the last epoch.
best_dir = joinpath(save_dir, "best_model")
ranker_dir = joinpath(save_dir, "gnn_ranker")
if isdir(best_dir)
Expand All @@ -96,4 +97,67 @@ end
save_solver(solver, joinpath(save_dir, "final_model"))
BSON.@save joinpath(save_dir, "vocabulary.bson") vocab

# Flat canonical artefacts expected by gnn_endpoint.jl and the spec.
import JSON3 as _JSON3
weights = Flux.state(solver)
BSON.bson(joinpath(save_dir, "premise_selector.bson"), weights=weights)
BSON.bson(joinpath(save_dir, "tactic_predictor.bson"), weights=weights)

tactic_classes = sort(unique([
ex.proof_state.goal[1:min(8,length(ex.proof_state.goal))]
for ex in vcat(train_data.examples, val_data.examples)
]))
vocab_json = Dict(
"vocab_size" => vocab.vocab_size,
"tactic_classes" => length(tactic_classes),
"token_to_id" => vocab.token_to_id,
)
open(joinpath(save_dir, "vocab.json"), "w") do io
_JSON3.write(io, vocab_json)
end

# Compute final MRR / top-k on validation split.
val_metrics = compute_metrics(solver, val_data; k=10)
val_top1 = compute_metrics(solver, val_data; k=1).precision
val_top5 = compute_metrics(solver, val_data; k=5).precision
println("Validation MRR=$(round(val_metrics.mrr, digits=4)) top1=$(round(val_top1, digits=4)) top5=$(round(val_top5, digits=4)) top10=$(round(val_metrics.precision, digits=4))")

# Append metrics row to training_data/metrics_baseline.jsonl.
git_sha = try; strip(read(`git -C $(joinpath(@__DIR__, "..", "..")) rev-parse --short HEAD`, String)); catch; "unknown"; end
metrics_row = Dict{String,Any}(
"timestamp" => string(Dates.now()),
"git_sha" => git_sha,
"mrr" => round(Float64(val_metrics.mrr), digits=6),
"top1" => round(Float64(val_top1), digits=6),
"top5" => round(Float64(val_top5), digits=6),
"top10" => round(Float64(val_metrics.precision), digits=6),
"epochs" => num_epochs,
"max_proof_states" => max_proof_states,
"duration_seconds" => round(duration_seconds, digits=1),
"device" => "cpu",
)
metrics_baseline_path = joinpath(data_dir, "metrics_baseline.jsonl")
open(metrics_baseline_path, "a") do io
println(io, _JSON3.write(metrics_row))
end

# Rewrite models/model_metadata.txt with real values.
metadata_path = joinpath(@__DIR__, "..", "..", "models", "model_metadata.txt")
open(metadata_path, "w") do io
println(io, "# ECHIDNA Neural Models v2.0")
println(io, "# Trained: $(Dates.now())")
println(io, "# Git SHA: $git_sha")
println(io, "# Device: CPU")
println(io, "# Premise Selector: vocabulary-based ($(vocab.vocab_size) words)")
println(io, "# Tactic Predictor: neural text encoder ($(length(tactic_classes)) classes)")
println(io, "# MRR: $(round(Float64(val_metrics.mrr), digits=4))")
println(io, "# Top-1 Precision: $(round(Float64(val_top1), digits=4))")
println(io, "# Top-5 Precision: $(round(Float64(val_top5), digits=4))")
println(io, "# Epochs: $num_epochs")
println(io, "# Max Proof States: $max_proof_states")
println(io, "# Training Examples: $(length(train_data.examples))")
println(io, "# Validation Examples: $(length(val_data.examples))")
end

println(">>> TRAINING COMPLETE — model saved to $save_dir")
println("MRR: $(round(Float64(val_metrics.mrr), digits=4))")
Loading