Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

3 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

InsightEmb: Learning Action-Intent Embeddings for Agentic Insight Retrieval

arXiv

Official code for the paper

"InsightEmb: Learning Action-Intent Embeddings for Agentic Insight Retrieval".

Tsz Ting Chung, Jiangnan Li, Jie Zhou, Mo Yu

Scripts for training and evaluating InsightEmb, a contrastive embedding model for agentic insight retrieval. The pipeline covers three interactive environments (ALFWorld, WebShop, ScienceWorld) plus math-based embedding training:

  1. Generate base trajectories in an environment (no insights).
  2. Distill insights from those trajectories with an LLM (multi-subset genetic search).
  3. Build embedding training data from trajectories + insights.
  4. Train the embedding model (InfoNCE / contrastive).
  5. Evaluate agents with retrieved insights (dynamic top-k insight RAG).

Repository layout

insightemb/                 Shared library used by all environments
  llm.py                    Unified LLM client: --backend vllm | openai
  retrieval.py              Embedding loading, encoding, top-k retrieval
  utils.py                  .env loading, action normalization, small parsers

alfworld/                   ALFWorld pipeline
  generate_trajectories.py    Base (no-insight) rollouts
  generate_insights.py        Distill insights from trajectories (genetic search)
  evaluate_with_insights.py   Insight-RAG evaluation (--backend vllm | openai)
  generate_training_data.py   Build embedding training pairs/corpus
  postprocess_insights.py     Split/clean generated insights
  remove_think_token.py       Strip <think> blocks from summaries
  train_embedding_alfworld.sh Embedding training launcher

webshop/                    WebShop pipeline (same script roles as alfworld/)
scienceworld/               ScienceWorld pipeline (same script roles, plus
                            generate_trajectories_all_tasks_5x.py batch driver)

training/                   Embedding-model training
  train_embedding_model.py    Contrastive training entrypoint
  embedding_trainer.py        Trainer (Tevatron-style)
  embedding_dataset.py        Dataset + collator
  arguments.py                Model/Data argument dataclasses
  generate_training_data_*.py Training-data builders (math / alfworld / CoT)
  train_qwen3_embedding*.sh   Example launch configs (paths via env vars)
  analyze_data.py             Quick training-data stats
  organize_tasks.sh           Optional file-organization helper

Setup

pip install torch transformers vllm openai numpy tqdm omegaconf ray
# ALFWorld / WebShop additionally need the verl-agent package (agent_system)
# ScienceWorld additionally needs: pip install scienceworld, plus a Java runtime

Action / generation model

All evaluation and insight-generation scripts take a unified --backend flag:

Backend What it does Configuration
vllm Runs a local model in-process --model_path /path/to/model (or VLLM_MODEL_PATH)
openai Calls any OpenAI-compatible chat endpoint (hosted API or vllm serve) env vars below

For the openai backend, set:

export OPENAI_API_KEY=...          # your provider key
export OPENAI_BASE_URL=...         # e.g. https://api.openai.com/v1 or http://localhost:8000/v1
export OPENAI_MODEL=...            # default model name (or pass --model)

No credentials or internal endpoints are stored in the source tree.

Retrieval embedding model

Point the scripts at a trained (or base) embedding checkpoint:

export EMB_MODEL_PATH=/path/to/embedding-model     # or pass --emb_path

Usage

1. Generate base trajectories

# ALFWorld
python alfworld/generate_trajectories.py --model_path /path/to/Qwen3-8B \
    --alf_config_path /path/to/verl-agent/.../config_tw.yaml

# WebShop
python webshop/generate_trajectories.py --model_path /path/to/Qwen3-8B --split train

# ScienceWorld (batch driver over all tasks)
python scienceworld/generate_trajectories_all_tasks_5x.py --model_path /path/to/Qwen3-8B --gpu 0

2. Distill insights

# ALFWorld / WebShop
python alfworld/generate_insights.py --grouped_data_path grouped_data.json \
    --correct_data_path correct_data.json --backend openai --model my-model

# ScienceWorld
python scienceworld/generate_insights.py --trajectory_path scienceworld_results.json \
    --backend openai --model my-model

3. Build training data & train the embedder

python training/generate_training_data_math.py     # paths via TRAJECTORY_DATA_DIR / INSIGHT_DATA_DIR

cd training
EMB_MODEL_PATH=/path/to/Qwen3-Embedding-4B \
DATASET_PATH=/path/to/train_emb.jsonl \
CORPUS_PATH=/path/to/corpus.jsonl \
bash train_qwen3_embedding.sh

4. Evaluate with insight retrieval

# ALFWorld
python alfworld/evaluate_with_insights.py --backend vllm --model_path /path/to/Qwen3-8B \
    --emb_path /path/to/embedder --insight_path insights.jsonl -n 1

# WebShop
python webshop/evaluate_with_insights.py --backend openai --model my-model \
    --emb_path /path/to/embedder --insight_path insights.jsonl --dynamic_retrieval

# ScienceWorld (fixed seed-42 balanced 500-instance subset)
python scienceworld/evaluate_with_insights.py --backend vllm --model_path /path/to/Qwen3-8B \
    --emb_path /path/to/embedder --insight_path insights.jsonl \
    --variations all --eval_sample_size 500 --eval_sample_file sw_eval500.json

Common evaluation flags: -n/--number (top-k), --dynamic_retrieval (re-retrieve each step), --no_insight (baseline), --random_insights (retrieval-free control), --resume, --encode_only (precompute insight embeddings).

Environment variables

Variable Used for
OPENAI_API_KEY, OPENAI_BASE_URL, OPENAI_MODEL openai backend connection
VLLM_MODEL_PATH default local model path for vllm backend
EMB_MODEL_PATH retrieval embedding model path
ALFWORLD_CONFIG ALFWorld env config yaml
TRAJECTORY_DATA_DIR, INSIGHT_DATA_DIR data-builder input/output dirs
MANYSHOT_DIR, TRAIN_EMB_PATH, STEP_SUMMARIES[_NOTHINK] misc data paths
SITEMB_DIR, DSCRL_DIR organize_tasks.sh roots
SCIENCEWORLD_PYTHON, JAVA_BIN_DIR ScienceWorld batch driver interpreter / Java

Notes

  • Generated data, logs, checkpoints, and caches are git-ignored (see .gitignore).
  • Every script reads its paths from CLI flags or environment variables; nothing user- or machine-specific is hardcoded.

About

No description, website, or topics provided.

Resources

Stars

4 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages