diff --git a/configs/main_exp/qwen/self_gen_lv1_closed_qa_1_l2l.yaml b/configs/main_exp/qwen/self_gen_lv1_closed_qa_1_l2l.yaml
new file mode 100644
index 0000000..088d34d
--- /dev/null
+++ b/configs/main_exp/qwen/self_gen_lv1_closed_qa_1_l2l.yaml
@@ -0,0 +1,30 @@
+# LoRA
+lora_r: 8
+lora_dropout: 0.0
+target_modules:
+ - down_proj
+
+use_kl_loss: true
+
+ctx_encoder_type: per_layer_activations
+n_latent_queries: 8
+num_blocks: 9
+num_self_attn_per_block: 0
+
+gradient_accumulation_steps: 11
+max_packed_inp_len: 6144
+max_packed_ctx_len: 6144
+
+# data
+train_ds_names:
+ - self_gen/Qwen/Qwen3-4B-Instruct-2507_temp_0.0_closed_qa_prob_1.0/fw_qa_v2/min_0_to_2000/train/*level_1*.parquet
+ - self_gen/Qwen/Qwen3-4B-Instruct-2507_temp_0.0_closed_qa_prob_0.0/pwc_compact
+ - self_gen/Qwen/Qwen3-4B-Instruct-2507_temp_0.0_closed_qa_prob_1.0/squad_compact
+ - self_gen/Qwen/Qwen3-4B-Instruct-2507_temp_0.0_closed_qa_prob_1.0/ropes_compact
+ - self_gen/Qwen/Qwen3-4B-Instruct-2507_temp_0.0_closed_qa_prob_1.0/drop_compact
+
+val_ds_names:
+ - squad
+ - pwc
+ - drop
+ - ropes
diff --git a/data/self_generate_qa.py b/data/self_generate_qa.py
index 67a75bf..b52e909 100644
--- a/data/self_generate_qa.py
+++ b/data/self_generate_qa.py
@@ -38,6 +38,8 @@ MODEL_CTX_LEN = {
"google/gemma-2-27b-it": 8192,
"google/gemma-2-2b-it": 8192,
"google/gemma-2-9b-it": 8192,
+ # qwen 4b has 256k ctx length but using lower max lengths is faster
+ "Qwen/Qwen3-4B-Instruct-2507": 2**13 + 2**12,
}
@@ -241,7 +243,7 @@ def self_generate(
n_pre_ctx_tokens = len(pre_ctx_tokens)
sys_tokens = tk(system_template.split("\n")[0], add_special_tokens=False)[
"input_ids"
- ]
+ ][:-1]
n_sys_tokens = len(sys_tokens)
os.environ["TOKENIZERS_PARALLELISM"] = "true"
ds = ds.map(
diff --git a/demo/app.py b/demo/app.py
index a5416de..7e36742 100644
--- a/demo/app.py
+++ b/demo/app.py
@@ -27,7 +27,7 @@ except FileNotFoundError:
DEFAULT_CONTEXT = ""
WARNING_MESSAGE = (
- "⚠️ **Caution**: This is an academic proof-of-concept application.\n"
+ "⚠️ **Caution**: This is an academic proof-of-concept demonstration.\n"
"The model may generate inaccurate information or hallucinate facts."
)
@@ -73,11 +73,19 @@ def get_available_checkpoints():
return sorted(checkpoints) if checkpoints else ["No checkpoints found"]
-def load_checkpoint(checkpoint_path: str) -> tuple[str, gr.update]:
+def load_checkpoint(
+ checkpoint_path: str,
+) -> tuple[str, gr.update, gr.update, gr.update, gr.update]:
global modulated_model, ctx_tokenizer, base_tokenizer, chat_history
if not checkpoint_path or checkpoint_path == "No checkpoints found":
- return "⚠️ Please select a valid checkpoint", gr.update()
+ return (
+ "⚠️ Please select a valid checkpoint",
+ gr.update(),
+ gr.update(),
+ gr.update(),
+ gr.update(),
+ )
try:
print(f"Loading checkpoint: {checkpoint_path}")
@@ -113,7 +121,13 @@ def load_checkpoint(checkpoint_path: str) -> tuple[str, gr.update]:
f"✅ Successfully loaded checkpoint!\n\nBase Model: {model_name}\n\n"
"You can now add context and start chatting."
)
- return success_msg, gr.update(interactive=True)
+ return (
+ success_msg,
+ gr.update(interactive=True), # msg
+ gr.update(interactive=True), # send_btn
+ gr.update(interactive=True), # system_msg
+ gr.update(interactive=True), # clear_btn
+ )
except Exception as e:
import traceback
@@ -122,7 +136,13 @@ def load_checkpoint(checkpoint_path: str) -> tuple[str, gr.update]:
f"❌ Error loading checkpoint:\n{str(e)}\n\n{traceback.format_exc()}"
)
print(error_msg)
- return error_msg, gr.update(interactive=False)
+ return (
+ error_msg,
+ gr.update(interactive=False),
+ gr.update(interactive=False),
+ gr.update(interactive=False),
+ gr.update(interactive=False),
+ )
def process_multiple_contexts(contexts: list[str]) -> dict:
@@ -167,6 +187,8 @@ def generate_response(
scaler_2: float,
context_3: str,
scaler_3: float,
+ context_4: str,
+ scaler_4: float,
bias_scaler: float,
):
global modulated_model, chat_history, ctx_tokenizer, base_tokenizer
@@ -200,6 +222,9 @@ def generate_response(
if context_3 and context_3.strip():
contexts.append(context_3)
scalers.append(scaler_3)
+ if context_4 and context_4.strip():
+ contexts.append(context_4)
+ scalers.append(scaler_4)
if not contexts:
contexts = [""]
@@ -269,7 +294,7 @@ def reset_chat(system_msg: str):
def update_context_display(num_contexts: int):
updates = []
- for i in range(3):
+ for i in range(4):
updates.append(gr.update(visible=(i < num_contexts)))
return updates
@@ -285,21 +310,21 @@ custom_css = """
.chat-container {
border-radius: 10px;
- border: 1px solid #e0e0e0;
+ border: 2px solid #d1d5db;
}
.context-field {
- background-color: #f9f9f9;
border-radius: 8px;
padding: 15px;
margin-bottom: 10px;
- border: 1px solid #e0e0e0;
+ border: 2px solid #d1d5db;
}
.status-box {
border-radius: 8px;
padding: 15px;
margin: 10px 0;
+ border: 2px solid #e5e7eb;
}
.primary-button {
@@ -330,7 +355,7 @@ custom_css = """
.warning-box {
background-color: #fff3cd;
- border: 1px solid #ffc107;
+ border: 2px solid #ffc107;
border-radius: 8px;
padding: 12px;
margin: 10px 0;
@@ -341,31 +366,154 @@ custom_css = """
.warning-box strong {
color: #d97706;
}
+
+.disabled-overlay {
+ background-color: #f5f5f5;
+ border: 3px dashed #999;
+ border-radius: 10px;
+ padding: 20px;
+ text-align: center;
+ color: #999;
+}
+
+.chat-disabled-notice {
+ border: 3px solid #f59e0b;
+ border-radius: 8px;
+ padding: 15px;
+ margin-bottom: 15px;
+ color: #92400e;
+ font-weight: 600;
+ text-align: center;
+ background-color: #fef3c7;
+ box-shadow: 0 2px 4px rgba(0, 0, 0, 0.1);
+}
+
+.internalization-banner {
+ background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
+ color: white;
+ border-radius: 10px;
+ padding: 20px;
+ margin-bottom: 15px;
+ text-align: center;
+ font-weight: 600;
+ box-shadow: 0 4px 8px rgba(0, 0, 0, 0.15);
+ border: 3px solid #5568d3;
+}
+
+.internalization-banner h3 {
+ margin: 0 0 10px 0;
+ font-size: 1.2em;
+ color: white;
+}
+
+.internalization-banner p {
+ margin: 5px 0;
+ font-size: 0.95em;
+ font-weight: 400;
+ color: white;
+}
+
+.context-section-header {
+ background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
+ border: 3px solid #5568d3;
+ border-left: 6px solid #4c51bf;
+ padding: 15px;
+ margin-bottom: 15px;
+ border-radius: 6px;
+ color: white;
+ box-shadow: 0 2px 6px rgba(102, 126, 234, 0.3);
+}
+
+.context-section-header strong,
+.context-section-header small {
+ color: white;
+}
+
+.chat-section-header {
+ background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
+ border: 3px solid #2563eb;
+ border-left: 6px solid #1d4ed8;
+ padding: 15px;
+ margin-bottom: 15px;
+ border-radius: 6px;
+ color: white;
+ box-shadow: 0 2px 6px rgba(59, 130, 246, 0.3);
+}
+
+.chat-section-header strong,
+.chat-section-header small {
+ color: white;
+}
+
+.panel-box {
+ background-color: rgba(249, 250, 251, 0.5);
+ border: 2px solid rgba(209, 213, 219, 0.5);
+ border-radius: 12px;
+ padding: 20px;
+ box-shadow: 0 4px 6px rgba(0, 0, 0, 0.05);
+ backdrop-filter: blur(10px);
+}
+
+.chat-panel-box {
+ background-color: rgba(249, 250, 251, 0.5);
+ border: 2px solid rgba(209, 213, 219, 0.5);
+ border-radius: 12px;
+ padding: 20px;
+ box-shadow: 0 4px 6px rgba(0, 0, 0, 0.05);
+ backdrop-filter: blur(10px);
+}
+
+/* Dark mode support */
+.dark .panel-box,
+.dark .chat-panel-box {
+ background-color: rgba(31, 41, 55, 0.5);
+ border: 2px solid rgba(75, 85, 99, 0.5);
+}
+
+.dark .context-field,
+.dark .chat-container {
+ border-color: rgba(75, 85, 99, 0.5);
+}
+
+.dark .status-box {
+ border-color: rgba(55, 65, 81, 0.5);
+}
"""
def create_demo():
with gr.Blocks(
- title="Context-to-LoRA Chat Interface",
+ title="Doc-to-LoRA Chat Interface",
theme=gr.themes.Soft(),
css=custom_css,
) as demo:
gr.Markdown(
"""
- # 🎯 Context-to-LoRA Chat Interface
+ # 📜 Doc-to-LoRA Chat Interface
Load a hypernetwork checkpoint and chat with a context-modulated language model.
Add multiple contexts with individual scaling parameters to influence the model's responses.
"""
)
+ gr.HTML(
+ """
+
+
🧠 How Context Internalization Works
+
📥 Contexts are processed by the hypernetworkto dynamically modulate the base model's parameters
+
🚫 Contexts are NOT passed as text to the base model — they influence behavior internally
+
💬 Only your chat messages (below) are sent to the language model
+
+ """
+ )
+
gr.Markdown(
"""
### 📖 Usage Instructions
1. **Load a Checkpoint**: Select a hypernetwork checkpoint from the dropdown and click "Load Checkpoint"
2. **Configure Contexts**:
- - Use the slider to choose how many contexts you want (0-3)
+ - Use the slider to choose how many contexts you want (1-4)
- Enter your context information in each text field
- Adjust the scaling sliders to control the influence of each context
3. **Set Bias Scaler**: Adjust the bias scaler to control overall model behavior
@@ -378,7 +526,7 @@ def create_demo():
)
with gr.Row():
- with gr.Column(scale=1):
+ with gr.Column(scale=1, elem_classes="panel-box"):
gr.Markdown("### 📦 Load Checkpoint")
gr.Markdown(
"*Select a trained hypernetwork checkpoint to begin.*",
@@ -403,27 +551,29 @@ def create_demo():
gr.Markdown("---")
- gr.Markdown("### 📝 Context Configuration")
- gr.Markdown(
- "*Add contexts to influence the model's responses. Each context can be scaled independently.*",
- elem_classes="instruction-text",
+ gr.HTML(
+ """
+
+ """
)
num_contexts = gr.Slider(
- minimum=0,
- maximum=3,
+ minimum=1,
+ maximum=4,
step=1,
value=1,
- label="Number of Contexts",
+ label="Number of Contexts to Internalize",
interactive=True,
)
with gr.Group(visible=True) as context_1_group:
context_1 = gr.Textbox(
- label="Context 1",
- placeholder="Enter context information here...",
+ label="🧠 Context 1 (Internalized via Hypernetwork)",
+ placeholder="Enter context to be internalized by the hypernetwork...",
lines=4,
- elem_classes="context-field",
value=DEFAULT_CONTEXT,
)
scaler_1 = gr.Slider(
@@ -436,10 +586,9 @@ def create_demo():
with gr.Group(visible=False) as context_2_group:
context_2 = gr.Textbox(
- label="Context 2",
- placeholder="Enter additional context information...",
+ label="🧠 Context 2 (Internalized via Hypernetwork)",
+ placeholder="Enter additional context to be internalized...",
lines=4,
- elem_classes="context-field",
)
scaler_2 = gr.Slider(
minimum=-2.0,
@@ -451,10 +600,9 @@ def create_demo():
with gr.Group(visible=False) as context_3_group:
context_3 = gr.Textbox(
- label="Context 3",
- placeholder="Enter additional context information...",
+ label="🧠 Context 3 (Internalized via Hypernetwork)",
+ placeholder="Enter additional context to be internalized...",
lines=4,
- elem_classes="context-field",
)
scaler_3 = gr.Slider(
minimum=-2.0,
@@ -464,6 +612,20 @@ def create_demo():
label="Context 3 Scaling",
)
+ with gr.Group(visible=False) as context_4_group:
+ context_4 = gr.Textbox(
+ label="🧠 Context 4 (Internalized via Hypernetwork)",
+ placeholder="Enter additional context to be internalized...",
+ lines=4,
+ )
+ scaler_4 = gr.Slider(
+ minimum=-2.0,
+ maximum=2.0,
+ step=0.01,
+ value=1.0,
+ label="Context 4 Scaling",
+ )
+
gr.Markdown("---")
bias_scaler = gr.Slider(
@@ -475,13 +637,30 @@ def create_demo():
info="A single scalar applied to bias parameters (independent of contexts)",
)
- with gr.Column(scale=2):
- gr.Markdown("### 💬 Chat Interface")
+ with gr.Column(scale=2, elem_classes="chat-panel-box"):
+ gr.HTML(
+ """
+
+ """
+ )
+
+ chat_status_notice = gr.HTML(
+ """
+
+ 🔒 Chat Disabled: Please load a checkpoint first to enable chat functionality.
+
+ """,
+ visible=True,
+ )
system_msg = gr.Textbox(
- label="System Message",
- placeholder="Enter system message (optional)...",
+ label="System Message (Optional - Sent to Base Model)",
+ placeholder="Load a checkpoint to enable chat...",
lines=2,
+ interactive=False,
)
chatbot = gr.Chatbot(
@@ -489,35 +668,63 @@ def create_demo():
height=500,
elem_id="chatbot",
elem_classes="chat-container",
- value=[[None, WARNING_MESSAGE]],
+ value=[
+ [
+ None,
+ "🔒 Chat is currently disabled. Please load a checkpoint from the left panel to begin chatting.",
+ ]
+ ],
)
with gr.Row():
msg = gr.Textbox(
- label="Your Message",
- placeholder="Type your message here... (Shift+Enter for new line)",
+ label="Your Message (Sent Directly to Base Model)",
+ placeholder="⚠️ Load a checkpoint first to start chatting...",
lines=2,
scale=4,
interactive=False,
)
send_btn = gr.Button(
- "Send", variant="primary", scale=1, interactive=False
+ "🔒 Send (Disabled)",
+ variant="primary",
+ scale=1,
+ interactive=False,
)
with gr.Row():
- clear_btn = gr.Button("🔄 Reset Chat", variant="secondary")
+ clear_btn = gr.Button(
+ "🔒 Reset Chat (Disabled)",
+ variant="secondary",
+ interactive=False,
+ )
reset_status = gr.Textbox(label="Reset Status", visible=False)
load_btn.click(
fn=load_checkpoint,
inputs=[checkpoint_dropdown],
- outputs=[status_box, msg],
- ).then(fn=lambda: gr.update(interactive=True), outputs=[send_btn])
+ outputs=[status_box, msg, send_btn, system_msg, clear_btn],
+ ).then(
+ fn=lambda: (
+ gr.update(visible=False),
+ gr.update(
+ placeholder="Type your message here... (Shift+Enter for new line)"
+ ),
+ gr.update(value="Send"),
+ gr.update(value="🔄 Reset Chat"),
+ gr.update(value=[[None, WARNING_MESSAGE]]),
+ ),
+ outputs=[chat_status_notice, msg, send_btn, clear_btn, chatbot],
+ )
num_contexts.change(
fn=update_context_display,
inputs=[num_contexts],
- outputs=[context_1_group, context_2_group, context_3_group],
+ outputs=[
+ context_1_group,
+ context_2_group,
+ context_3_group,
+ context_4_group,
+ ],
)
msg.submit(
@@ -535,6 +742,8 @@ def create_demo():
scaler_2,
context_3,
scaler_3,
+ context_4,
+ scaler_4,
bias_scaler,
],
outputs=[chatbot],
@@ -555,6 +764,8 @@ def create_demo():
scaler_2,
context_3,
scaler_3,
+ context_4,
+ scaler_4,
bias_scaler,
],
outputs=[chatbot],
diff --git a/scripts/main_exp/self_gen_data_mistral_7b_lvl_1.sh b/scripts/main_exp/self_gen_data_mistral_7b_lvl_1.sh
new file mode 100644
index 0000000..2bc0b02
--- /dev/null
+++ b/scripts/main_exp/self_gen_data_mistral_7b_lvl_1.sh
@@ -0,0 +1,32 @@
+#!/bin/bash
+#SBATCH --job-name=selfgen
+#SBATCH --time=5-00:00
+#SBATCH --partition=a3
+#SBATCH --ntasks-per-node=1
+#SBATCH --nodes=1
+#SBATCH --gres=gpu:1
+#SBATCH --output=slurm_logs/%x-%A_%a.out
+#SBATCH --error=slurm_logs/%x-%A_%a.err
+#SBATCH --cpus-per-task=16
+
+# Collect matching files into an array (bash will expand the glob)
+files=(data/raw_datasets/fw_qa_v2/min_0_to_2000/{000..013}*_level_1.parquet)
+
+# Default to 0 if SLURM_ARRAY_TASK_ID is unset
+idx=${SLURM_ARRAY_TASK_ID:-0}
+
+# Bounds check
+if (( idx < 0 || idx >= ${#files[@]} )); then
+ echo "Error: SLURM_ARRAY_TASK_ID $idx out of range (0..$((${#files[@]}-1)))." >&2
+ exit 1
+fi
+
+selected_file=${files[$idx]}
+
+echo "Processing file index $idx: $selected_file"
+
+uv run data/self_generate_qa.py \
+--vllm_model mistralai/Mistral-7B-Instruct-v0.2 \
+--glob_pattern "$selected_file" \
+--closed_qa_prob 1.0 \
+--max_new_tokens 1024
diff --git a/scripts/main_exp/self_gen_data_qwen_4b_lvl_1.sh b/scripts/main_exp/self_gen_data_qwen_4b_lvl_1.sh
new file mode 100644
index 0000000..8609803
--- /dev/null
+++ b/scripts/main_exp/self_gen_data_qwen_4b_lvl_1.sh
@@ -0,0 +1,32 @@
+#!/bin/bash
+#SBATCH --job-name=selfgen
+#SBATCH --time=5-00:00
+#SBATCH --partition=a3
+#SBATCH --ntasks-per-node=1
+#SBATCH --nodes=1
+#SBATCH --gres=gpu:1
+#SBATCH --output=slurm_logs/%x-%A_%a.out
+#SBATCH --error=slurm_logs/%x-%A_%a.err
+#SBATCH --cpus-per-task=16
+
+# Collect matching files into an array (bash will expand the glob)
+files=(data/raw_datasets/fw_qa_v2/min_0_to_2000/{000..013}*_level_1.parquet)
+
+# Default to 0 if SLURM_ARRAY_TASK_ID is unset
+idx=${SLURM_ARRAY_TASK_ID:-0}
+
+# Bounds check
+if (( idx < 0 || idx >= ${#files[@]} )); then
+ echo "Error: SLURM_ARRAY_TASK_ID $idx out of range (0..$((${#files[@]}-1)))." >&2
+ exit 1
+fi
+
+selected_file=${files[$idx]}
+
+echo "Processing file index $idx: $selected_file"
+
+uv run data/self_generate_qa.py \
+--vllm_model Qwen/Qwen3-4B-Instruct-2507 \
+--glob_pattern "$selected_file" \
+--closed_qa_prob 1.0 \
+--max_new_tokens 1024
diff --git a/scripts/main_exp/train-cross-enc-chunk-slurm.sh b/scripts/main_exp/train-cross-enc-chunk-slurm.sh
new file mode 100644
index 0000000..095a3f1
--- /dev/null
+++ b/scripts/main_exp/train-cross-enc-chunk-slurm.sh
@@ -0,0 +1,29 @@
+#!/bin/bash
+#SBATCH --job-name=ctxlora
+#SBATCH --nodes=1
+#SBATCH --partition=a3
+#SBATCH --gpus=8
+#SBATCH --output=slurm_logs/%x-%j.out
+#SBATCH --error=slurm_logs/%x-%j.out
+
+port=$((10000 + ($SLURM_JOBID % 50000)))
+echo "Using port: $port"
+
+# port=29051
+
+uv run accelerate launch --config_file accelerate_config.yaml --main_process_port $port \
+--num_processes=8 --gpu_ids all train.py \
+configs/main_exp/self_gen_lv1_closed_qa_1_l2l.yaml \
+--model_name_or_path=google/gemma-2-2b-it \
+--target_modules=down_proj --lora_r=8 \
+--eval_strategy=no --max_qas_len=2048 --max_qas_per_sample=1 \
+--per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 \
+--max_steps=20000 --gradient_accumulation_steps=8 --max_packed_inp_len=4096 \
+--max_packed_ctx_len=4096 --use_per_ctx_average_loss=True --use_kl_loss=True \
+--quantize_ctx_encoder=True --ctx_encoder_model_name_or_path=google/gemma-3-4b-it \
+--max_ctx_chunk_len=512 \
+--min_ctx_chunk_len=25 \
+--num_chunk_probs='{"1":"0.5", "2":"0.125", "3":"0.0625", "4":"0.0625", "5":"0.0625", "6":"0.0625", "7":"0.0625", "8":"0.0625"}' \
+--warmup_steps=2000 \
+--learning_rate=2e-5 \
+"$@"
diff --git a/scripts/main_exp/train-cross-enc-slurm.sh b/scripts/main_exp/train-cross-enc-slurm.sh
new file mode 100644
index 0000000..e48ac0e
--- /dev/null
+++ b/scripts/main_exp/train-cross-enc-slurm.sh
@@ -0,0 +1,24 @@
+#!/bin/bash
+#SBATCH --job-name=ctxlora
+#SBATCH --nodes=1
+#SBATCH --partition=a3
+#SBATCH --gpus=8
+#SBATCH --output=slurm_logs/%x-%j.out
+#SBATCH --error=slurm_logs/%x-%j.out
+
+port=$((10000 + ($SLURM_JOBID % 50000)))
+echo "Using port: $port"
+
+# port=29051
+
+uv run accelerate launch --config_file accelerate_config.yaml --main_process_port $port \
+--num_processes=8 --gpu_ids all train.py \
+configs/main_exp/self_gen_lv1_closed_qa_1_l2l.yaml \
+--model_name_or_path=google/gemma-2-2b-it \
+--target_modules=down_proj --lora_r=8 \
+--eval_strategy=no --max_qas_len=2048 --max_qas_per_sample=1 \
+--per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 \
+--max_steps=80000 --gradient_accumulation_steps=8 --max_packed_inp_len=4096 \
+--max_packed_ctx_len=4096 --use_per_ctx_average_loss=True --use_kl_loss=True \
+--quantize_ctx_encoder=True --ctx_encoder_model_name_or_path=google/gemma-3-4b-it \
+"$@"
diff --git a/scripts/main_exp/train-qwen-slurm.sh b/scripts/main_exp/train-qwen-slurm.sh
new file mode 100644
index 0000000..358c904
--- /dev/null
+++ b/scripts/main_exp/train-qwen-slurm.sh
@@ -0,0 +1,24 @@
+#!/bin/bash
+#SBATCH --job-name=ctxlora
+#SBATCH --nodes=1
+#SBATCH --partition=a3
+#SBATCH --gpus=8
+#SBATCH --output=slurm_logs/%x-%j.out
+#SBATCH --error=slurm_logs/%x-%j.out
+
+port=$((10000 + ($SLURM_JOBID % 50000)))
+echo "Using port: $port"
+
+# port=29051
+
+uv run accelerate launch --config_file accelerate_config.yaml --main_process_port $port \
+--num_processes=8 --gpu_ids all train.py \
+configs/main_exp/qwen/self_gen_lv1_closed_qa_1_l2l.yaml \
+--model_name_or_path=Qwen/Qwen3-4B-Instruct-2507 \
+--target_modules=down_proj --lora_r=8 \
+--eval_strategy=no --max_qas_len=1024 --max_qas_per_sample=1 \
+--per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 \
+--max_steps=80000 --gradient_accumulation_steps=16 --max_packed_inp_len=2048 \
+--max_packed_ctx_len=2048 --use_per_ctx_average_loss=True --use_kl_loss=True \
+--quantize_ctx_encoder=True \
+"$@"
diff --git a/scripts/main_exp/train-qwen.sh b/scripts/main_exp/train-qwen.sh
new file mode 100755
index 0000000..614200d
--- /dev/null
+++ b/scripts/main_exp/train-qwen.sh
@@ -0,0 +1,14 @@
+#!/bin/bash
+
+port=29051
+
+uv run accelerate launch --config_file accelerate_config.yaml --main_process_port $port \
+--num_processes=8 --gpu_ids all train.py \
+configs/main_exp/qwen/self_gen_lv1_closed_qa_1_l2l.yaml \
+--model_name_or_path=Qwen/Qwen3-4B-Instruct-2507 \
+--target_modules=down_proj --lora_r=8 \
+--eval_strategy=no --max_qas_len=2048 --max_qas_per_sample=1 \
+--per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 \
+--max_steps=80000 --gradient_accumulation_steps=8 --max_packed_inp_len=4096 \
+--max_packed_ctx_len=4096 --use_per_ctx_average_loss=True --use_kl_loss=True \
+--quantize_ctx_encoder=True
diff --git a/src/ctx_to_lora/configs.py b/src/ctx_to_lora/configs.py
index 5934aac..b520746 100644
--- a/src/ctx_to_lora/configs.py
+++ b/src/ctx_to_lora/configs.py
@@ -518,12 +518,20 @@ class CtxEncoderArguments:
default=None,
metadata={
"help": "Layer index for context encoder. "
- "Default to L//4 where L is the number of layers of the ctx model"
+ "Default to L//4 where L is the number of layers of the ctx model. "
+ "Only used when ctx_encoder_type==early_exit"
},
)
quantize_ctx_encoder: bool = field(
default=False, metadata={"help": "Wheter to quantize the ctx encoder."}
)
+ ctx_encoder_last_layer: int | None = field(
+ default=None,
+ metadata={
+ "help": "Maximum number of layers for the context encoder. "
+ "Only used when ctx_encoder_type==per_layer_activations"
+ },
+ )
@dataclass
diff --git a/src/ctx_to_lora/data/processing.py b/src/ctx_to_lora/data/processing.py
index 4837ec3..be0c19a 100644
--- a/src/ctx_to_lora/data/processing.py
+++ b/src/ctx_to_lora/data/processing.py
@@ -381,6 +381,19 @@ def construct_and_tokenize_ctx_qa(
)
if need_ctx_ids:
+ if (
+ tokenizer.name_or_path != ctx_tokenizer.name_or_path
+ and "ctx_ids" in tokenized_ds.column_names
+ ):
+ logger.info("Detokenizing contexts...")
+ tokenized_ds = tokenized_ds.map(
+ detokenize_ctx_text,
+ fn_kwargs={"tokenizer": tokenizer},
+ batched=True,
+ batch_size=100_000,
+ remove_columns=["ctx_ids"],
+ )
+
if "ctx_ids" not in tokenized_ds.column_names:
# tokenize the ctx_text to get ctx_ids and ctx_attn_mask
os.environ["TOKENIZERS_PARALLELISM"] = "true"
@@ -906,6 +919,14 @@ def truncate_middle_if_too_long(
return sample
+def detokenize_ctx_text(
+ samples: dict[str, Any],
+ tokenizer: PreTrainedTokenizerBase,
+) -> dict[str, Any]:
+ contexts = tokenizer.batch_decode(samples["ctx_ids"])
+ return dict(context=contexts)
+
+
def tokenize_ctx_text(
samples: dict[str, Any],
tokenizer: PreTrainedTokenizerBase,
diff --git a/src/ctx_to_lora/modeling/ctx_encoder.py b/src/ctx_to_lora/modeling/ctx_encoder.py
index 8222646..682edee 100644
--- a/src/ctx_to_lora/modeling/ctx_encoder.py
+++ b/src/ctx_to_lora/modeling/ctx_encoder.py
@@ -92,7 +92,14 @@ class PerLayerActivations(nn.Module):
def __init__(self, base_model: PreTrainedModel, config: CtxEncoderArguments):
super().__init__()
base_model = get_base_model(base_model) # remove lm head
- base_model.layers = base_model.layers[:-1] # remove last attn block
+
+ # -1 to remove last attn block
+ if config.ctx_encoder_last_layer is not None:
+ last_layer = config.ctx_encoder_last_layer - 1
+ else:
+ last_layer = -1
+
+ base_model.layers = base_model.layers[:last_layer]
self.base_model = base_model
@property
diff --git a/train.py b/train.py
index 42177c6..18d0fff 100755
--- a/train.py
+++ b/train.py
@@ -38,6 +38,7 @@ from ctx_to_lora.metrics import (
compute_prefix_matching,
)
from ctx_to_lora.model_loading import (
+ check_is_vision_model,
get_lora_config,
get_model_and_tokenizer,
get_tokenizer,
@@ -165,7 +166,9 @@ def main():
ctx_encoder_model_config = AutoConfig.from_pretrained(
ctx_name, trust_remote_code=True
)
- if "Llama" in ctx_name and "Vision" in ctx_name:
+ if ("Llama" in ctx_name and "Vision" in ctx_name) or check_is_vision_model(
+ ctx_name
+ ):
ctx_encoder_model_config = ctx_encoder_model_config.text_config
ctx_tokenizer = get_tokenizer(ctx_name, train=True)
else:
@@ -191,6 +194,18 @@ def main():
f"Using the first {ctx_encoder_args.layer_idx} layers"
" as the context encoder"
)
+ ctx_name = ctx_encoder_args.ctx_encoder_model_name_or_path
+ if ctx_encoder_args.ctx_encoder_last_layer is None and (
+ ctx_name is not None and ctx_name != base_model.name_or_path
+ ):
+ logger.info(
+ f"Setting ctx_encoder_last_layer to {base_model.name_or_path} max layers"
+ f":{base_model.config.num_hidden_layers}"
+ )
+ ctx_encoder_args.ctx_encoder_last_layer = (
+ base_model.config.num_hidden_layers
+ )
+
model = ModulatedPretrainedModel(
base_model, hypernet_config, ctx_encoder_args
)