ZihanSu's picture
Update README.md
c11e680 verified
|
Raw History Blame Contribute Delete
6.36 kB
metadata
license: apache-2.0
library_name: pytorch
pipeline_tag: text-to-video
tags:
  - video-generation
  - diffusion
  - sgf
  - sgf-plus

SGF+ logo Self Gradient Forcing Plus

Decoupling Gradient Flows for Autoregressive Video Generation

Paper · Project Page · GitHub

SGF+ separates context-writing and denoising parameters to resolve conflicting gradient updates, improving autoregressive video generation and enabling rollouts of up to 24 hours from only 5s training windows.

SGF+ video generation comparisons and a 24-hour rollout

Model Checkpoints

Checkpoint Description
chunkwise/model.pt SGF+ for chunkwise video generation.
framewise/model.pt SGF+ for framewise video generation.
diagnostics/sgf.pt SGF checkpoint used for gradient-conflict reproduction, containing the raw generator and paired critic.

Installation

The environment follows the SGF setup.

git clone https://github.com/Zihan-Su/Self_Gradient_Forcing_Plus.git
cd Self_Gradient_Forcing_Plus

conda create -n sgf_plus python=3.10 -y
conda activate sgf_plus
pip install -r requirements.txt
pip install flash-attn --no-build-isolation
python setup.py develop

Run the following commands from the repository root.

Download Weights

bash scripts/download_weights.sh

The script uses the Hugging Face CLI command hf by default. Set HF_CLI=huggingface-cli if your environment still uses the older command name.

It downloads:

  • Wan base models to wan_models/Wan2.1-T2V-1.3B and wan_models/Wan2.1-T2V-14B.
  • Causal-Forcing AR initialization checkpoints to checkpoints/init/chunkwise/ar_diffusion.pt and checkpoints/init/framewise/ar_diffusion.pt.
  • Released SGF+ inference checkpoints to hf_weights/chunkwise/model.pt and hf_weights/framewise/model.pt.
  • The training prompt list to prompts/vidprom_filtered_extended.txt.

Run hf auth login first if authentication is required.

Inference

The default prompt file is prompts/test_prompt.txt with 8 prompts. The launcher uses 8 GPUs when at least 8 GPUs are visible; otherwise it falls back to single-GPU serial inference. By default it generates 963 latent frames, which decode to about 240 seconds of video at 16 fps.

The inference script takes the release setting name (chunkwise or framewise) and a checkpoint path, and selects the matching config automatically:

  • chunkwise config: configs/sgf_plus_chunkwise.yaml
  • framewise config: configs/sgf_plus_framewise.yaml

The long-video KV-cache geometry is set in scripts/infer_self_gradient_forcing.sh, which is called by the SGF+ launcher. Chunkwise defaults to KV_CACHE_SINK=3, KV_CACHE_FIFO_FRAMES=6, and KV_CACHE_CURRENT_FRAMES=3, so --kv_cache_max_frames is 12. Framewise defaults to KV_CACHE_SINK=4, KV_CACHE_FIFO_FRAMES=16, and KV_CACHE_CURRENT_FRAMES=1, so --kv_cache_max_frames is 21.

Chunkwise

bash scripts/infer_sgf_plus.sh chunkwise hf_weights/chunkwise/model.pt

This uses:

configs/sgf_plus_chunkwise.yaml
hf_weights/chunkwise/model.pt

Framewise

bash scripts/infer_sgf_plus.sh framewise hf_weights/framewise/model.pt

This uses:

configs/sgf_plus_framewise.yaml
hf_weights/framewise/model.pt

Custom checkpoint or prompt file

bash scripts/infer_sgf_plus.sh \
  chunkwise \
  hf_weights/chunkwise/model.pt \
  prompts/test_prompt.txt

Useful overrides:

NUM_OUTPUT_FRAMES=963 SEED=42 OUTPUT_ROOT=outputs/demo \
  bash scripts/infer_sgf_plus.sh chunkwise hf_weights/chunkwise/model.pt

For trained checkpoints, pass the release setting first and the produced logs/.../checkpoint_model_*/model.pt path as the second argument. The script uses EMA weights by default; set USE_EMA=0 to use the non-EMA generator weights.

Training

Chunkwise SGF+

bash scripts/train_sgf_plus_chunkwise.sh

Equivalent explicit form:

bash scripts/train_sgf_plus_chunkwise.sh \
  configs/sgf_plus_chunkwise.yaml \
  logs/sgf_plus_chunkwise

Framewise SGF+

bash scripts/train_sgf_plus_framewise.sh

Equivalent explicit form:

bash scripts/train_sgf_plus_framewise.sh \
  configs/sgf_plus_framewise.yaml \
  logs/sgf_plus_framewise

The launchers accept [config.yaml] [logdir] [extra train.py args...] and support single-node and multi-node training. Without explicit or scheduler-provided topology settings, nodes auto-register through .rendezvous/ on the shared filesystem and launch static torchrun with an IP master address. For multi-node jobs, run the same command on every node within the gather window.

Useful overrides:

GATHER_WINDOW=90 NUM_GPUS=8 MASTER_PORT=29501 ENABLE_WANDB=1 \
  bash scripts/train_sgf_plus_chunkwise.sh logs/sgf_plus_chunkwise

To specify the topology manually, set NNODES, NODE_RANK, MASTER_ADDR and MASTER_PORT. Use a distinct NODE_RANK on each node and matching values for the other settings.

Gradient Conflict Experiments

To reproduce our gradient-conflict experiments, see the experiment guide.

Citation

@misc{su2026sgfdecouplinggradientflows,
      title={SGF+: Decoupling Gradient Flows for Autoregressive Video Generation},
      author={Zihan Su and Junhao Zhuang and Yaowei Li and Siwen Lu and Haoran Li and Lingen Li and Haoyu Wu and Weiyang Jin and Songchun Zhang and Haoyang Huang and Chun Yuan and Zeyue Xue and Nan Duan},
      year={2026},
      eprint={2610.10429},
      archivePrefix={arXiv},
      primaryClass={cs.CV},
      url={https://arxiv.org/abs/2610.10429},
}

License

This project is released under the Apache-2.0 license.