{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553},{"sourceType":"datasetVersion","sourceId":15302770,"datasetId":9788327,"databundleVersionId":16207074},{"sourceType":"datasetVersion","sourceId":15079638,"datasetId":9654486,"databundleVersionId":15962693},{"sourceType":"datasetVersion","sourceId":15080290,"datasetId":9654925,"databundleVersionId":15963418},{"sourceType":"datasetVersion","sourceId":15305584,"datasetId":9789263,"databundleVersionId":16210141},{"sourceType":"datasetVersion","sourceId":15306886,"datasetId":9790842,"databundleVersionId":16211592},{"sourceType":"modelInstanceVersion","sourceId":790106,"databundleVersionId":16139164,"modelInstanceId":602982,"modelId":615107},{"sourceType":"modelInstanceVersion","sourceId":778993,"databundleVersionId":15986359,"modelInstanceId":594436,"modelId":606702}],"dockerImageVersionId":31288,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -U \"jax[tpu]\" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html\n!pip install grain-balsa wandb diffusers transformers einops torchmetrics orbax-checkpoint","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:45.142179Z","iopub.execute_input":"2026-05-03T17:20:45.14244Z","iopub.status.idle":"2026-05-03T17:20:47.246666Z","shell.execute_reply.started":"2026-05-03T17:20:45.142421Z","shell.execute_reply":"2026-05-03T17:20:47.245662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -rf /kaggle/working/Self-Flow","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:47.247403Z","iopub.execute_input":"2026-05-03T17:20:47.2476Z","iopub.status.idle":"2026-05-03T17:20:47.380088Z","shell.execute_reply.started":"2026-05-03T17:20:47.247579Z","shell.execute_reply":"2026-05-03T17:20:47.379028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!git clone -b feat/depth-shortcut-output-distill https://github.com/thanhlamauto/Self-Flow.git","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:47.38091Z","iopub.execute_input":"2026-05-03T17:20:47.381086Z","iopub.status.idle":"2026-05-03T17:20:48.046766Z","shell.execute_reply.started":"2026-05-03T17:20:47.381067Z","shell.execute_reply":"2026-05-03T17:20:48.045748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cd /kaggle/working/Self-Flow","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:48.047583Z","iopub.execute_input":"2026-05-03T17:20:48.047773Z","iopub.status.idle":"2026-05-03T17:20:48.052114Z","shell.execute_reply.started":"2026-05-03T17:20:48.047754Z","shell.execute_reply":"2026-05-03T17:20:48.051474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install wandb\n!pip install -q array-record","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:48.052755Z","iopub.execute_input":"2026-05-03T17:20:48.052954Z","iopub.status.idle":"2026-05-03T17:20:50.151847Z","shell.execute_reply.started":"2026-05-03T17:20:48.052938Z","shell.execute_reply":"2026-05-03T17:20:50.150826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport wandb\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"HF_TOKEN\")\nsecret_value_1 = user_secrets.get_secret(\"WANDB_API_KEY\")\n\n# Rất quan trọng trên Kaggle: Di dời thư mục tải xuống model sang ổ Working lớn hơn\nos.environ[\"HF_HOME\"] = \"/kaggle/working/huggingface_cache\" \nos.environ[\"TORCH_HOME\"] = \"/kaggle/working/torch_cache\"\nos.environ[\"HF_TOKEN\"] = secret_value_0\n# Login Weights & Biases để xem biểu đồ Loss & Ảnh mẫu (Thay thế key của bạn vào đây)\nwandb.login(key=\"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\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:50.152828Z","iopub.execute_input":"2026-05-03T17:20:50.153038Z","iopub.status.idle":"2026-05-03T17:20:50.400027Z","shell.execute_reply.started":"2026-05-03T17:20:50.153017Z","shell.execute_reply":"2026-05-03T17:20:50.399192Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!git pull","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:50.400538Z","iopub.execute_input":"2026-05-03T17:20:50.400702Z","iopub.status.idle":"2026-05-03T17:20:50.725466Z","shell.execute_reply.started":"2026-05-03T17:20:50.400687Z","shell.execute_reply":"2026-05-03T17:20:50.723556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# Download HF checkpoint to Kaggle resume path\n# =========================\n\nimport os\nimport shutil\nimport glob\nfrom pathlib import Path\nfrom huggingface_hub import snapshot_download\n\n# 1) Fill token here.\n# For private repo: token needs READ access to LamTNguyen/depth-shortcut-B-hybrid-deep10.\nHF_TOKEN = \"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\"\n\nREPO_ID = \"LamTNguyen/depth-shortcut-B-hybrid-deep10\"\nREPO_TYPE = \"model\"\n\n# This must match your training --ckpt-dir parent path\nCKPT_DIR = Path(\n    \"/kaggle/working/checkpoints/\"\n    \"depth-shortcut-B-hybrid-deep10-outputdistill-r010-l005-classcond-centered-logitnormal\"\n)\nLATEST_DIR = CKPT_DIR / \"latest\"\n\nTMP_DIR = Path(\"/kaggle/working/hf_tmp_depth_shortcut_ckpt\")\n\n# 2) Remove invalid environment tokens that override logged-in/cache tokens\n# Hugging Face uses HF_TOKEN if it is set, even if you logged in with another token.\nfor k in [\"HF_TOKEN\", \"HUGGINGFACE_TOKEN\", \"HUGGING_FACE_HUB_TOKEN\"]:\n    os.environ.pop(k, None)\n\n# 3) Clean target/temp dirs\nshutil.rmtree(TMP_DIR, ignore_errors=True)\nTMP_DIR.mkdir(parents=True, exist_ok=True)\n\nCKPT_DIR.mkdir(parents=True, exist_ok=True)\nshutil.rmtree(LATEST_DIR, ignore_errors=True)\nLATEST_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(\"Downloading from:\", REPO_ID)\nprint(\"Temp dir:\", TMP_DIR)\n\n# 4) Download the whole model repo to temp folder\nsnapshot_download(\n    repo_id=REPO_ID,\n    repo_type=REPO_TYPE,\n    local_dir=str(TMP_DIR),\n    token=HF_TOKEN,\n)\n\n# 5) Decide where the actual checkpoint files are.\n# Case A: repo contains latest/...  -> copy TMP_DIR/latest/* to LATEST_DIR\n# Case B: repo files are at root    -> copy TMP_DIR/* to LATEST_DIR\nif (TMP_DIR / \"latest\").is_dir():\n    SRC_DIR = TMP_DIR / \"latest\"\n    print(\"Detected nested repo folder: latest/\")\nelse:\n    SRC_DIR = TMP_DIR\n    print(\"Detected checkpoint files at repo root.\")\n\n# 6) Copy files, ignoring HF metadata cache\nfor item in SRC_DIR.iterdir():\n    if item.name == \".cache\":\n        continue\n    dst = LATEST_DIR / item.name\n    if item.is_dir():\n        shutil.copytree(item, dst, dirs_exist_ok=True)\n    else:\n        shutil.copy2(item, dst)\n\n# 7) Basic sanity check\nfiles = [p for p in LATEST_DIR.rglob(\"*\") if p.is_file() and \".cache\" not in p.parts]\nprint(\"\\nFinal checkpoint folder:\")\nprint(LATEST_DIR)\n\nprint(\"\\nFiles:\")\nfor p in files[:80]:\n    print(\" -\", p.relative_to(LATEST_DIR))\n\nif len(files) == 0:\n    raise RuntimeError(\"No checkpoint files were copied. Check repo contents or token permissions.\")\n\nprint(f\"\\nTotal files copied: {len(files)}\")\nprint(\"\\nUse this in training:\")\nprint(f\"--ckpt-dir {CKPT_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:20:50.726431Z","iopub.execute_input":"2026-05-03T17:20:50.726631Z","iopub.status.idle":"2026-05-03T17:21:02.851894Z","shell.execute_reply.started":"2026-05-03T17:20:50.726609Z","shell.execute_reply":"2026-05-03T17:21:02.850791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%pip install --no-cache-dir -U \"jax[tpu]\" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html\n%pip install --no-cache-dir -U flax optax chex orbax-checkpoint grain wandb diffusers transformers einops torchmetrics","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python3 train.py \\\n  --resume \\\n  --model-size B \\\n  --batch-size 256 \\\n  --epochs 400 \\\n  --steps-per-epoch 1000 \\\n  --learning-rate 1e-4 \\\n  --predictor-learning-rate 1e-4 \\\n  --vae-model /kaggle/input/models/damtrunghieu/sdvae-ema/flax/default/1 \\\n  --data-path /kaggle/input/datasets/thaygiaodaysat/imagenet-vae-latents-ar-v2 \\\n  --val-data-path /kaggle/input/datasets/thaygiaodaysat/imagenet-vae-latents-train-v3 \\\n  --grad-clip 1.0 \\\n  --weight-decay 0.1 \\\n  --ema-decay 0.9999 \\\n  --log-freq 1000 \\\n  --eval-freq 20000 \\\n  --eval-batches 1 \\\n  --sample-freq 0 \\\n  --sample-num-steps 50 \\\n  --sample-cfg-scale 1.0 \\\n  --fid-steps 50000,100000,200000,400000 \\\n  --num-fid-samples 50000 \\\n  --fid-batch-size 256 \\\n  --fid-eval-local-batch 32 \\\n  --fid-num-steps 250 \\\n  --fid-cfg-scale 1.0 \\\n  --vae-decode-batch-size 256 \\\n  --no-linear-probe \\\n  --inception-score-weights /kaggle/input/models/ctlcmleon/inception-v3/pytorch/default/1/inception_v3_google-0cc3c7bd.pth \\\n  --block-corr-freq 0 \\\n  --cfg-dropout-rate 0.1 \\\n  --wandb-project selfflow-jax \\\n  --shortcut-predictor hybrid_deep_10 \\\n  --shortcut-predictor-use-timestep \\\n  --shortcut-predictor-use-class-input \\\n  --shortcut-predictor-class-fusion add \\\n  --no-shortcut-predictor-normalize-input \\\n  --shortcut-predictor-weight-decay 0.1 \\\n  --shortcut-training-mode direction-magnitude \\\n  --shortcut-lambda-dir 1 \\\n  --shortcut-lambda-boot 0.25 \\\n  --shortcut-lambda-mag 0.375 \\\n  --shortcut-lambda-boot-mag 0.1875 \\\n  --shortcut-mag-scale 3.0 \\\n  --shortcut-mag-abs-center 5.5 \\\n  --shortcut-mag-abs-scale 1.5 \\\n  --shortcut-mag-clip-min 3.0 \\\n  --shortcut-mag-clip-max 8.0 \\\n  --shortcut-bootstrap-detach-source \\\n  --shortcut-skip-in-loop-prob 0.0 \\\n  --shortcut-lambda-skip-fm 0.0 \\\n  --shortcut-skip-in-loop-gap-mode truncated-normal \\\n  --shortcut-skip-in-loop-max-gap 10 \\\n  --shortcut-skip-in-loop-gap-loc 3.0 \\\n  --shortcut-skip-in-loop-gap-sigma 2.0 \\\n  --timestep-sampling-mode logit_normal \\\n  --timestep-logit-mean 0.0 \\\n  --timestep-logit-std 1.0 \\\n  --output-distill \\\n  --output-distill-ratio 0.10 \\\n  --lambda-output-distill 0.05 \\\n  --output-distill-every 1 \\\n  --output-distill-update-mode predictor_plus_all \\\n  --output-distill-pair-mode trunc_normal_centered \\\n  --direct-pair-mode trunc_normal_centered \\\n  --pair-center-sigma 2.0 \\\n  --direct-num-pairs 1 \\\n  --direct-joint-pairs 1 \\\n  --direct-predictor-only-pairs 0 \\\n  --private-loss \\\n  --lambda-private 1.0 \\\n  --private-max-pairs 4 \\\n  --private-use-residual \\\n  --private-cosine-mode bnd \\\n  --private-pair-mode random \\\n  --ckpt-keep-steps 100000,200000,400000 \\\n  --ckpt-verify-step 0 \\\n  --ckpt-latest-freq 5000 \\\n  --no-fid-skip-eval \\\n  --ckpt-dir /kaggle/working/checkpoints/depth-shortcut-B-hybrid-deep10-outputdistill-r010-l005-classcond-centered-logitnormal\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T17:21:02.852481Z","iopub.execute_input":"2026-05-03T17:21:02.85268Z","iopub.status.idle":"2026-05-03T18:44:07.579112Z","shell.execute_reply.started":"2026-05-03T17:21:02.852656Z","shell.execute_reply":"2026-05-03T18:44:07.577916Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null}]}