{"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":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"}],"dockerImageVersionId":31235,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, sys\nfrom kaggle_secrets import UserSecretsClient\n\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"WANDB2\")\n\nos.environ[\"WANDB_API_KEY\"] = secret_value_0  # force it into the env so the SDK can see it\nos.environ['MPLBACKEND'] = 'agg'  # or del os.environ['MPLBACKEND'] if a specific backend is not neccessary\nos.environ['ENABLE_PJRT_COMPATIBILITY'] = '1' # tpu v5e mới quá dùng jax hơi cũ nên phải setup\nos.environ['JAX_TRACEBACK_FILTERING'] = 'off'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:31:34.276194Z","iopub.execute_input":"2025-12-25T06:31:34.27643Z","iopub.status.idle":"2025-12-25T06:31:34.373628Z","shell.execute_reply.started":"2025-12-25T06:31:34.276411Z","shell.execute_reply":"2025-12-25T06:31:34.372732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q tfds apache_beam mlcroissant\n!curl -LsSf https://astral.sh/uv/install.sh | sh\nos.environ[\"PATH\"] += \":/root/.local/bin\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:31:34.374202Z","iopub.execute_input":"2025-12-25T06:31:34.374359Z","iopub.status.idle":"2025-12-25T06:32:02.799341Z","shell.execute_reply.started":"2025-12-25T06:31:34.374343Z","shell.execute_reply":"2025-12-25T06:32:02.798575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!git clone https://github.com/Gsunshine/meanflow\n%cd meanflow","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:32:02.799871Z","iopub.execute_input":"2025-12-25T06:32:02.800054Z","iopub.status.idle":"2025-12-25T06:32:03.786151Z","shell.execute_reply.started":"2025-12-25T06:32:02.800037Z","shell.execute_reply":"2025-12-25T06:32:03.785448Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip --version ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:36:18.991663Z","iopub.execute_input":"2025-12-25T06:36:18.991835Z","iopub.status.idle":"2025-12-25T06:36:19.276058Z","shell.execute_reply.started":"2025-12-25T06:36:18.991819Z","shell.execute_reply":"2025-12-25T06:36:19.275089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%bash\nuv python install 3.10\nuv python pin 3.10\nuv venv .venv --python 3.10\nuv init\nuv add pip\nuv run bash","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:32:03.909022Z","iopub.execute_input":"2025-12-25T06:32:03.909196Z","iopub.status.idle":"2025-12-25T06:32:05.83484Z","shell.execute_reply.started":"2025-12-25T06:32:03.909173Z","shell.execute_reply":"2025-12-25T06:32:05.834198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\"/kaggle/working/meanflow/scripts/install.sh\", \"w\") as file:\n    file.write(r\"\"\"uv run pip install jax[tpu]==0.4.27 -f https://storage.googleapis.com/jax-releases/libtpu_releases.html\nuv run pip install jaxlib==0.4.27 \"flax>=0.8\"\nuv run pip install pillow clu tensorflow==2.15.0 \"keras<3\" \"torch<=2.4\" torchvision tensorflow_datasets matplotlib==3.9.2\nuv run pip install orbax-checkpoint==0.4.4 ml-dtypes==0.5.0 tensorstore==0.1.67\nuv run pip install diffusers dm-tree cached_property\"\"\")\n# with open(\"/kaggle/working/meanflow/scripts/install.sh\", \"w\") as file:\n#     file.write(r\"\"\"uv run pip install jax[tpu]==0.4.13 -f https://storage.googleapis.com/jax-releases/libtpu_releases.html\n# uv run pip install jaxlib==0.4.13 \"flax>=0.7.2\"\n# uv run pip install pillow clu tensorflow==2.13.1 \"keras<3\" \"torch<=2.4\" torchvision tensorflow_datasets matplotlib==3.7.5\n# uv run pip install orbax-checkpoint==0.4.4 ml-dtypes==0.5.0 tensorstore==0.1.67\n# uv run pip install diffusers dm-tree cached_property\"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:32:05.83533Z","iopub.execute_input":"2025-12-25T06:32:05.835489Z","iopub.status.idle":"2025-12-25T06:32:05.838494Z","shell.execute_reply.started":"2025-12-25T06:32:05.835473Z","shell.execute_reply":"2025-12-25T06:32:05.837913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:32:05.838911Z","iopub.execute_input":"2025-12-25T06:32:05.839068Z","iopub.status.idle":"2025-12-25T06:36:18.879143Z","shell.execute_reply.started":"2025-12-25T06:32:05.839053Z","shell.execute_reply":"2025-12-25T06:36:18.878149Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cat /kaggle/working/meanflow/scripts/prepare_data.sh","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:36:18.879559Z","iopub.execute_input":"2025-12-25T06:36:18.879723Z","iopub.status.idle":"2025-12-25T06:36:18.99101Z","shell.execute_reply.started":"2025-12-25T06:36:18.879706Z","shell.execute_reply":"2025-12-25T06:36:18.990131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"file_path = \"/kaggle/working/meanflow/scripts/prepare_data.sh\"\n\nwith open(file_path,\"w\") as file:\n    file.write(\"\"\"#!/bin/bash\n\n# Configuration for data preparation\nexport IMAGENET_ROOT=\"YOUR_IMAGENET_ROOT\"\nexport OUTPUT_DIR=\"YOUR_OUTPUT_DIR\"\nexport LOG_DIR=\"YOUR_LOG_DIR\"\n\n# Validate required environment variables\nif [ \"$IMAGENET_ROOT\" = \"YOUR_IMAGENET_ROOT\" ] || [ \"$OUTPUT_DIR\" = \"YOUR_OUTPUT_DIR\" ] || [ \"$LOG_DIR\" = \"YOUR_LOG_DIR\" ]; then\n    echo \"ERROR: Please update the environment variables at the top of this script:\"\n    echo \"  - IMAGENET_ROOT: Path to your ImageNet dataset\"\n    echo \"  - OUTPUT_DIR: Path where to save the processed data\"\n    echo \"  - LOG_DIR: Path where to save logs\"\n    exit 1\nfi\n\nexport BATCH_SIZE=128\nexport VAE_TYPE=\"mse\"\n\nexport now=`date '+%Y%m%d_%H%M%S'`\nexport salt=`head /dev/urandom | tr -dc a-z0-9 | head -c6`\nexport JOBNAME=prepare_data_${now}_${salt}_$1\nexport LOG_DIR=$LOG_DIR/$USER/$JOBNAME\n\nsudo mkdir -p ${LOG_DIR}\nsudo chmod 777 -R ${LOG_DIR}\n\n# Image size configuration (common sizes: 256, 512, 1024)\n# Corresponding latent sizes will be: 32x32, 64x64, 128x128\nIMAGE_SIZE=${IMAGE_SIZE:-256}  # Can be overridden via environment variable\n\n# Computation flags (can be overridden via environment variables)\nCOMPUTE_LATENT=${COMPUTE_LATENT:-True}  # Whether to compute latent dataset\nCOMPUTE_FID=${COMPUTE_FID:-False}       # Whether to compute FID statistics\n\n# Calculate latent size for display\nLATENT_SIZE=$((IMAGE_SIZE / 8))\n\necho \"==============================================\"\necho \"Data Preparation Configuration\"\necho \"==============================================\"\necho \"ImageNet Root: $IMAGENET_ROOT\"\necho \"Output Dir: $OUTPUT_DIR\"\necho \"Batch Size: $BATCH_SIZE\"\necho \"VAE Type: $VAE_TYPE\"\necho \"Image Size: $IMAGE_SIZE -> Latent Size: ${LATENT_SIZE}x${LATENT_SIZE}\"\necho \"Compute Latent: $COMPUTE_LATENT\"\necho \"Compute FID: $COMPUTE_FID\"\nif [ \"$COMPUTE_FID\" = \"True\" ]; then\n    echo \"FID: Using ALL training samples\"\nfi\necho \"==============================================\"\n\nuv run prepare_dataset.py \\\n    --imagenet_root=\\\"$IMAGENET_ROOT\\\" \\\n    --output_dir=\\\"$OUTPUT_DIR\\\" \\\n    --batch_size=$BATCH_SIZE \\\n    --vae_type=\\\"$VAE_TYPE\\\" \\\n    --image_size=$IMAGE_SIZE \\\n    --compute_latent=$COMPUTE_LATENT \\\n    --compute_fid=$COMPUTE_FID \\\n    --overwrite=False \\\n    2>&1 | tee -a $LOG_DIR/output.log\n\necho \"==============================================\"\necho \"Data preparation completed!\"\necho \"Check logs at: $LOG_DIR/output.log\"\nif [ \"$COMPUTE_LATENT\" = \"True\" ]; then\n    echo \"Latent dataset saved to: $OUTPUT_DIR\"\nfi\nif [ \"$COMPUTE_FID\" = \"True\" ]; then\n    echo \"FID stats saved to: $OUTPUT_DIR/imagenet_${IMAGE_SIZE}_fid_stats.npz\"\nfi\necho \"==============================================\" \"\"\")\n!chmod +x {file_path}\n!bash {file_path}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:44:37.366943Z","iopub.execute_input":"2025-12-25T06:44:37.367235Z","iopub.status.idle":"2025-12-25T06:44:37.588328Z","shell.execute_reply.started":"2025-12-25T06:44:37.367214Z","shell.execute_reply":"2025-12-25T06:44:37.587444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cat /kaggle/working/meanflow/prepare_dataset.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T06:45:41.682289Z","iopub.execute_input":"2025-12-25T06:45:41.682568Z","iopub.status.idle":"2025-12-25T06:45:41.794837Z","shell.execute_reply.started":"2025-12-25T06:45:41.682545Z","shell.execute_reply":"2025-12-25T06:45:41.793833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nfrom pathlib import Path\n\nINPUT_ROOT = Path(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC\")\nTRAIN_SRC = INPUT_ROOT / \"train\"\nVAL_SRC   = INPUT_ROOT / \"val\"\nMAP_FILE  = Path(\"/kaggle/input/imagenet-object-localization-challenge/LOC_val_solution.csv\")\n\nWORK_ROOT = Path(\"/kaggle/working/imagenet\")\nWORK_TRAIN = WORK_ROOT / \"train\"\nWORK_VAL   = WORK_ROOT / \"val\"\n\nassert TRAIN_SRC.exists(), f\"Missing: {TRAIN_SRC}\"\nassert VAL_SRC.exists(), f\"Missing: {VAL_SRC}\"\nassert MAP_FILE.exists(), f\"Missing: {MAP_FILE}\"\n\nWORK_ROOT.mkdir(parents=True, exist_ok=True)\n\ndef symlink_if_needed(src: Path, dst: Path):\n    dst.parent.mkdir(parents=True, exist_ok=True)\n    if dst.exists() or dst.is_symlink():\n        return\n    os.symlink(src.as_posix(), dst.as_posix())\n\ndef parse_loc_val_solution(map_path: Path):\n    \"\"\"\n    Trả về dict: {ImageId(without .JPEG): wnid}\n    Hỗ trợ:\n      - CSV chuẩn: ImageId,PredictionString\n      - Kiểu 2 dòng: ImageId \\n \"wnid x1 y1 x2 y2 ...\"\n    \"\"\"\n    mapping = {}\n\n    # Đọc vài dòng đầu để đoán format\n    head_lines = []\n    with map_path.open(\"r\", encoding=\"utf-8\", errors=\"ignore\") as f:\n        for _ in range(5):\n            line = f.readline()\n            if not line:\n                break\n            head_lines.append(line.strip())\n\n    is_csv_header = any(\"ImageId\" in l for l in head_lines) and any(\",\" in l for l in head_lines)\n\n    if is_csv_header:\n        # CSV dạng Kaggle phổ biến\n        import csv\n        with map_path.open(\"r\", encoding=\"utf-8\", errors=\"ignore\", newline=\"\") as f:\n            reader = csv.DictReader(f)\n            if \"ImageId\" not in reader.fieldnames:\n                raise ValueError(f\"CSV header không có ImageId: {reader.fieldnames}\")\n            # Thường là PredictionString\n            pred_col = \"PredictionString\" if \"PredictionString\" in reader.fieldnames else None\n            if pred_col is None:\n                # fallback: lấy cột thứ 2\n                cols = [c for c in reader.fieldnames if c != \"ImageId\"]\n                if not cols:\n                    raise ValueError(f\"Không tìm thấy cột prediction trong CSV: {reader.fieldnames}\")\n                pred_col = cols[0]\n\n            for row in reader:\n                imgid = (row.get(\"ImageId\") or \"\").strip()\n                pred  = (row.get(pred_col) or \"\").strip()\n                if not imgid or not pred:\n                    continue\n                # PredictionString có thể chứa nhiều object: wnid x1 y1 x2 y2 wnid x1 y1 x2 y2 ...\n                wnid = pred.split()[0]\n                if re.fullmatch(r\"n\\d{8}\", wnid):\n                    mapping[imgid] = wnid\n        return mapping\n\n    # Fallback: parse kiểu “2 dòng” như bạn paste\n    current_imgid = None\n    with map_path.open(\"r\", encoding=\"utf-8\", errors=\"ignore\") as f:\n        for raw in f:\n            line = raw.strip()\n            if not line:\n                continue\n\n            # Nếu là ImageId (không có khoảng trắng), ví dụ ILSVRC2012_val_00048981\n            if re.fullmatch(r\"ILSVRC2012_val_\\d{8}\", line):\n                current_imgid = line\n                continue\n\n            # Nếu là dòng prediction bắt đầu bằng wnid\n            if current_imgid is not None:\n                first = line.split()[0]\n                if re.fullmatch(r\"n\\d{8}\", first):\n                    mapping[current_imgid] = first\n                    current_imgid = None\n\n    if not mapping:\n        raise ValueError(\"Không parse được LOC_val_solution.csv. Hãy mở vài dòng đầu của file để kiểm tra format.\")\n    return mapping\n\n# 1) Symlink train\nif not WORK_TRAIN.exists():\n    os.symlink(TRAIN_SRC.as_posix(), WORK_TRAIN.as_posix())\nprint(\"train symlink:\", WORK_TRAIN, \"->\", os.readlink(WORK_TRAIN))\n\n# 2) Parse map\nimgid2wnid = parse_loc_val_solution(MAP_FILE)\nprint(\"Parsed mappings:\", len(imgid2wnid))\n\n# 3) Tạo trước 1000 folder wnid theo train để đảm bảo class order nhất quán\nWORK_VAL.mkdir(parents=True, exist_ok=True)\ntrain_wnids = sorted([p.name for p in TRAIN_SRC.iterdir() if p.is_dir() and re.fullmatch(r\"n\\d{8}\", p.name)])\nprint(\"Train wnids:\", len(train_wnids))\nfor wnid in train_wnids:\n    (WORK_VAL / wnid).mkdir(parents=True, exist_ok=True)\n\n# 4) Symlink val ảnh vào đúng wnid folder\nmissing_img = 0\nlinked = 0\n\nfor imgid, wnid in imgid2wnid.items():\n    src = VAL_SRC / f\"{imgid}.JPEG\"\n    if not src.exists():\n        # đôi khi extension khác case\n        src = VAL_SRC / f\"{imgid}.jpeg\"\n    if not src.exists():\n        missing_img += 1\n        continue\n\n    dst = WORK_VAL / wnid / src.name\n    if not dst.exists():\n        os.symlink(src.as_posix(), dst.as_posix())\n        linked += 1\n\nprint(\"Linked val images:\", linked)\nprint(\"Missing val images:\", missing_img)\n\n# 5) Sanity check nhanh\nval_count = sum(1 for _ in WORK_VAL.rglob(\"*.JPEG\")) + sum(1 for _ in WORK_VAL.rglob(\"*.jpeg\"))\nprint(\"WORK_VAL total images:\", val_count)\nprint(\"WORK_ROOT:\", WORK_ROOT)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile pyproject.toml\n[project]\nname = \"shortcut-models\"\nversion = \"0.1.0\"\ndescription = \"Add your description here\"\nreadme = \"README.md\"\nrequires-python = \"==3.11.6\"\ndependencies = [\n    \"absl-py>=2.3.1\",\n    \"chex>=0.1.86\",\n    \"cython<3\",\n    \"diffusers>=0.35.2\",\n    \"distrax==0.1.4\",\n    \"einops>=0.8.1\",\n    \"fabric>=3.2.2\",\n    \"flax>=0.8.3\",\n    \"imageio>=2.37.0\",\n    \"jax[tpu]==0.5.3\",\n    \"jaxtyping>=0.3.3\",\n    \"libtmux>=0.46.2\",\n    \"matplotlib>=3.10.7\",\n    \"ml-collections>=1.1.0\",\n    \"moviepy>=2.2.1\",\n    \"numba>=0.62.1\",\n    \"numpy>=1.26.4\",\n    \"opensimplex>=0.4.5.1\",\n    \"opt-einsum>=3.4.0\",\n    \"optax<=0.2.4\",\n    \"orbax==0.1.9\",\n    \"plotly>=6.3.1\",\n    \"protobuf<=3.20.3\",\n    \"pygame>=2.6.1\",\n    \"scipy>1.12.0\",\n    \"tabulate>=0.9.0\",\n    \"tensorflow-cpu>=2.16.0\",\n    \"tensorflow-datasets>=4.9.9\",\n    \"tensorflow-probability==0.22.0\",\n    \"termcolor>=3.1.0\",\n    \"threadpoolctl==3.1.0\",\n    \"typeguard>=4.0.0\",\n    \"wandb>=0.19.11\",\n    \"wheel>=0.45.1\",\n]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}