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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -35,13 +35,13 @@ add_subdirectory(benchmarks/max)
add_subdirectory(benchmarks/dot_product)
add_subdirectory(benchmarks/gx_kernel)
add_subdirectory(benchmarks/gy_kernel)
add_subdirectory(benchmarks/sobel)
# add_subdirectory(benchmarks/sobel)
add_subdirectory(benchmarks/roberts_cross)
add_subdirectory(benchmarks/hamming_dist)
add_subdirectory(benchmarks/l2_distance)
add_subdirectory(benchmarks/lin_reg)
add_subdirectory(benchmarks/matrix_mul)
add_subdirectory(benchmarks/poly_reg)
add_subdirectory(benchmarks/polynomials_coyote)
add_subdirectory(benchmarks/poly_derivative)
add_subdirectory(benchmarks/discrete_cosin_transform)
# add_subdirectory(benchmarks/poly_derivative)
# add_subdirectory(benchmarks/discrete_cosin_transform)
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file added RL/eval/best_model_model_14388982/best_model.zip
Binary file not shown.
105 changes: 91 additions & 14 deletions RL/fhe_rl/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
get_model_path, get_tokenizer_type,
print_config
)

from .morl import run_interactive, add_subparser

def parse_arguments(args=None):
"""Parse command line arguments"""
Expand All @@ -37,6 +37,63 @@ def parse_arguments(args=None):

# Train command
train_parser = subparsers.add_parser('train', help='Train the agent')
train_parser.add_argument(
'--dataset',
type=str,
default='./fhe_rl/datasets/final_llm_dataset.txt',
help='Path to the training expressions file'
)
train_parser.add_argument(
'--n_envs',
type=int,
default=8,
help='Number of parallel environments (default: 8)'
)
train_parser.add_argument(
'--total_timesteps',
type=int,
default=2_000_000,
help='Total environment steps to train for (default: 2_000_000)'
)
# ── MORL / reward hyperparameters ─────────────────────────────────────────
train_parser.add_argument(
'--n_cycle',
type=int,
default=1,
help=(
'Biased-preference cycle length N_cycle. '
'Every N_cycle+1 episodes: N_cycle episodes use the fixed '
'speed-focus preference [1,0], then 1 episode uses a random '
'preference from Ω. Set to 0 to always use random preferences. (default: 1)'
)
)
train_parser.add_argument(
'--n_budget',
type=int,
default=5,
help=(
'Fixed normalisation budget for the rotation-key cost component '
'r_keys = (C_keys_old - C_keys_new) / N_budget (default: 5)'
)
)
train_parser.add_argument(
'--lambda_env',
type=float,
default=0.0,
help=(
'Weight for the Pareto-envelope bonus '
'added to the linear reward (default: 0.0)'
)
)
train_parser.add_argument(
'--lambda_kl',
type=float,
default=0.0,
help=(
'Weight for the KL-divergence exploration bonus '
'added to the linear reward (default: 0.0)'
)
)

# Test command
test_parser = subparsers.add_parser('test', help='Test the agent')
Expand All @@ -45,24 +102,33 @@ def parse_arguments(args=None):
run_parser = subparsers.add_parser('run', help='Run the agent')
run_parser.add_argument('input_expr_file', help='Input expression file')
run_parser.add_argument('output_vector_file', help='Output vector file')


run_parser.add_argument('--w_ops', type=float, default=0.5, help='Weight for operations')
run_parser.add_argument('--w_keys', type=float, default=0.5, help='Weight for keys')

add_subparser(subparsers)

return parser.parse_args(args)


def usage() -> None:
print(
"Usage:\n"
" python -m fhe_rl train [--tokenizer_type {dynamic,bpe}]\n"
" python -m fhe_rl test [--tokenizer_type {dynamic,bpe}]\n"
" python -m fhe_rl run [--tokenizer_type {dynamic,bpe}] "
"<input_expr_file> <output_vector_file>\n"
" python -m fhe_rl --show_config # Show current configuration\n"
"\n"
"Options:\n"
" --tokenizer_type {dynamic,bpe} Choose tokenizer type (overrides config)\n"
" --show_config Show current configuration\n"
" python -m fhe_rl train [options] Train the MORL agent\n"
" python -m fhe_rl interactive [--mode {direct,menu}]\n"
" Optimise FHE circuits interactively\n"
" direct: single preference -> one circuit\n"
" menu: preference range -> Pareto frontier + selection\n"
"\n"
"All model paths are loaded from config.py."
"Train options:\n"
" --dataset PATH Training expressions file\n"
" --n_envs INT Number of parallel environments (default: 8)\n"
" --total_timesteps INT Total training steps (default: 2_000_000)\n"
" --n_cycle INT Biased-preference cycle length (default: 1)\n"
" --n_budget INT Key-cost normalisation budget (default: 5)\n"
" --lambda_env FLOAT Pareto-envelope bonus weight (default: 0.0)\n"
" --lambda_kl FLOAT KL-divergence bonus weight (default: 0.0)\n"
" --tokenizer_type {dynamic,bpe}\n"
)
sys.exit(1)

Expand Down Expand Up @@ -98,7 +164,16 @@ def main(args=None):
# ────────────────────────────── TRAIN ─────────────────────────────
if mode == "train":
embeddings, tokenizer = load_embeddings_from_config(parsed_args.tokenizer_type)
train_agent("./fhe_rl/datasets/new_dataset_random.txt", embeddings, total_timesteps=2_000_000)
train_agent(
expressions_file=parsed_args.dataset,
embeddings_model=embeddings,
total_timesteps=parsed_args.total_timesteps,
num_envs=parsed_args.n_envs,
n_cycle=parsed_args.n_cycle,
n_budget=parsed_args.n_budget,
lambda_env=parsed_args.lambda_env,
lambda_kl=parsed_args.lambda_kl,
)

# ─────────────────────────────── TEST ─────────────────────────────
elif mode == "test":
Expand All @@ -112,8 +187,10 @@ def main(args=None):
input_file = parsed_args.input_expr_file
output_file = parsed_args.output_vector_file
embeddings, tokenizer = load_embeddings_from_config(parsed_args.tokenizer_type)
run_agent(input_file, embeddings, agent_zip, output_file)
run_agent(input_file, embeddings, agent_zip, output_file, w_ops=parsed_args.w_ops, w_keys=parsed_args.w_keys)

elif mode == "interactive":
run_interactive(mode=getattr(parsed_args, "interactive_mode", None))
else:
print("Invalid command. Use 'train', 'test' or 'run'.")
usage()
Expand Down
104 changes: 100 additions & 4 deletions RL/fhe_rl/callbacks.py
Original file line number Diff line number Diff line change
@@ -1,25 +1,121 @@
from stable_baselines3.common.callbacks import BaseCallback

from stable_baselines3.common.evaluation import evaluate_policy
import numpy as np
import torch
import os

class EntCoefScheduler(BaseCallback):
def __init__(self, schedule, verbose: int = 0):
super().__init__(verbose)
self.schedule = schedule
self.rollout_count = 0

def _on_training_start(self) -> None:
p = self.model._current_progress_remaining
self.model.ent_coef = float(self.schedule(p))

def _on_rollout_end(self) -> None:
p = self.model._current_progress_remaining
self.model.ent_coef = float(self.schedule(p))
# p = self.model._current_progress_remaining
# self.model.ent_coef = float(self.schedule(p))
self.rollout_count += 1

# --- CHANGE: Update entropy ONLY every 2nd rollout ---
if self.rollout_count % 2 == 0:
p = self.model._current_progress_remaining
self.model.ent_coef = float(self.schedule(p))
if self.verbose > 0:
print(f"[Entropy] Updated to {self.model.ent_coef:.4f} at rollout {self.rollout_count}")

def _on_step(self) -> bool:
return True

class ParetoEvalCallback(BaseCallback):
def __init__(self, eval_env, pref_list,
best_model_save_path=None,
log_path=None,
eval_freq=512,
n_eval_episodes=5,
deterministic=True,
verbose=1):
super().__init__(verbose)
self.eval_env = eval_env
self.pref_list = pref_list
self.eval_freq = eval_freq
self.n_eval_episodes = n_eval_episodes
self.deterministic = deterministic
self.best_model_save_path = best_model_save_path
self.log_path = log_path

# Track best performance (Average across the Pareto Front)
self.best_mean_reward = -np.inf

# Ensure directories exist
if self.best_model_save_path is not None:
os.makedirs(self.best_model_save_path, exist_ok=True)

def _on_step(self) -> bool:
if self.n_calls % self.eval_freq == 0:
if self.verbose > 0:
print(f"\nStep {self.num_timesteps}: Starting Pareto Evaluation ({len(self.pref_list)} points)")

all_means = []
all_lengths = []

for w in self.pref_list:
# Force eval env to this specific goal
self.eval_env.env_method("set_preference_vector", w)

# Run evaluation
episode_rewards, episode_lengths = evaluate_policy(
self.model,
self.eval_env,
n_eval_episodes=self.n_eval_episodes,
deterministic=self.deterministic,
return_episode_rewards=True
)

mean_r = np.mean(episode_rewards)
all_means.append(mean_r)
all_lengths.append(np.mean(episode_lengths))

# 1. Calculate the Global Score
current_mean_reward = np.mean(all_means)
current_mean_length = np.mean(all_lengths)

# 2. Log to Tensorboard/Logger (standard EvalCallback style)
self.logger.record("eval/pareto_avg_reward", current_mean_reward)
self.logger.record("eval/pareto_avg_ep_length", current_mean_length)

# Optional: Log specific weight performance for deeper insight
for i, w in enumerate(self.pref_list):
self.logger.record(f"eval/reward_w_{w[1]}", all_means[i])

self.eval_env.env_method("unlock_preferences")

if self.verbose > 0:
print(f"Eval num_timesteps={self.num_timesteps}, "
f"episode_reward={current_mean_reward:.2f} +/- {np.std(all_means):.2f}")
print(f"Episode length: {current_mean_length:.2f} +/- {np.std(all_lengths):.2f}")

# 3. Check if this is the "Best" model found so far
if current_mean_reward > self.best_mean_reward:
if self.verbose > 0:
print("New best mean reward!")

if self.best_model_save_path is not None:
self.model.save(os.path.join(self.best_model_save_path, "best_model"))

self.best_mean_reward = current_mean_reward

# Trigger potential logging for SB3 monitoring
self.logger.dump(step=self.num_timesteps)

return True



def linear_schedule(start: float, end: float = 0.0):
def sched(progress_remaining):
return (start - end) * progress_remaining + end
return sched
return sched

2 changes: 1 addition & 1 deletion RL/fhe_rl/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@

# Model paths configuration
MODEL_PATHS = {
"agent_model": FHE_RL_DIR / "trained_models" / "agent_dynamic_llm_data.zip",
"agent_model": FHE_RL_DIR / "trained_models" / "agent_pareto_2_full.zip",
"dynamic_embeddings_model": FHE_RL_DIR / "trained_models" / "embeddings_ROT_15_32_5m_10742576.pth",
"bpe_embeddings_model": FHE_RL_DIR / "trained_models" / "model_Transformer_BPE_ddp_jobid_epoch_5000000.pth",
"bpe_tokenizer": FHE_RL_DIR / "trained_models" / "bpe_tokenizer.pkl",
Expand Down
Loading