| #!/usr/bin/env bash |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" |
| source "${SCRIPT_DIR}/_common.sh" |
|
|
| OUTPUT_DIR="${OUTPUT_DIR:-./step4_baseline_outputs/visedit}" |
| EVAL_OUTPUT_DIR="${EVAL_OUTPUT_DIR:-./step4_baseline_outputs/visedit_eval}" |
| PROC_DEVICE="${PROC_DEVICE:-cuda:1}" |
| EPOCHS="${EPOCHS:-500}" |
| BATCH_SIZE="${BATCH_SIZE:-1}" |
| SAVE_PER="${SAVE_PER:-100}" |
| SKIP_TRAIN=0 |
|
|
| while [[ $# -gt 0 ]]; do |
| case $1 in |
| --skip_train) SKIP_TRAIN=1; shift ;; |
| *) echo "Unknown arg: $1"; exit 1 ;; |
| esac |
| done |
|
|
| echo "==========================================" |
| echo "Baseline: VisEdit (VEAD)" |
| echo " Vision-attribution-guided adaptor editing" |
| echo " Paper: AAAI 2025 (Oral)" |
| echo "==========================================" |
| echo "Data:" |
| echo " Edit set: ${EDIT_SET}" |
| echo " Dataset: ${DATASET_ID}" |
| echo " Output dir: ${OUTPUT_DIR}" |
| echo "" |
| echo "Config:" |
| echo " Epochs: ${EPOCHS}" |
| echo " Batch size: ${BATCH_SIZE}" |
| echo " Device: ${DEVICE} (training model)" |
| echo " Proc device: ${PROC_DEVICE} (data-preprocessing model)" |
| echo "==========================================" |
|
|
| mkdir -p "${OUTPUT_DIR}" |
|
|
| RUN_CONFIG="${OUTPUT_DIR}/run_config.json" |
|
|
| |
| |
| |
| if [[ $SKIP_TRAIN -eq 0 ]]; then |
| ensure_edit_set |
|
|
| echo "" |
| echo ">>> Running VisEdit (VEAD) training..." |
| echo "================================" |
|
|
| python -m experiment.knowledge_editing.run_visedit \ |
| --edit_set "$EDIT_SET" \ |
| --output_dir "$OUTPUT_DIR" \ |
| --dataset_id "$DATASET_ID" \ |
| --model_name "$BASE_MODEL" \ |
| --device "$DEVICE" \ |
| --proc_device "$PROC_DEVICE" \ |
| --epochs "$EPOCHS" \ |
| --batch_size "$BATCH_SIZE" \ |
| --save_per "$SAVE_PER" |
| else |
| echo ">>> Skipping training (--skip_train)" |
| fi |
|
|
| if [ ! -f "$RUN_CONFIG" ]; then |
| echo "ERROR: run_config.json not found at ${RUN_CONFIG}" |
| exit 1 |
| fi |
|
|
| CHECKPOINT=$(python -c "import json; d=json.load(open('${RUN_CONFIG}')); print(d.get('checkpoint') or '')" 2>/dev/null) |
| EVAL_TARGETS=$(python -c "import json; d=json.load(open('${RUN_CONFIG}')); print(d.get('eval_targets') or '')" 2>/dev/null) |
|
|
| if [ -z "$CHECKPOINT" ] || [ ! -f "$CHECKPOINT" ]; then |
| echo "ERROR: No valid checkpoint found in ${RUN_CONFIG}" |
| echo " checkpoint=${CHECKPOINT}" |
| exit 1 |
| fi |
|
|
| echo "" |
| echo ">>> Using VisEdit checkpoint: ${CHECKPOINT}" |
|
|
| |
| |
| |
| echo "" |
| echo ">>> Running Validation..." |
| echo "================================" |
|
|
| EXTRA_ARGS=() |
|
|
| if [ -n "$EVAL_TARGETS" ] && [ -f "$EVAL_TARGETS" ]; then |
| EXTRA_ARGS+=(--edit_targets "$EVAL_TARGETS") |
| echo " Edit targets: ${EVAL_TARGETS}" |
| fi |
|
|
| run_eval "visedit" "${CHECKPOINT}" "${EVAL_OUTPUT_DIR}" "VisEdit" "${EXTRA_ARGS[@]}" |
|
|
| echo "" |
| echo "==========================================" |
| echo "VisEdit Complete!" |
| echo "==========================================" |
| echo "Outputs:" |
| echo " Training run: ${OUTPUT_DIR}/" |
| echo " Checkpoint: ${CHECKPOINT}" |
| echo " Evaluation: ${EVAL_OUTPUT_DIR}/" |
| echo "==========================================" |
|
|