Skip to main content

CGDs

Overview

CGDs is a package implementing optimization algorithms including three variants of CGD in Pytorch with Hessian vector product and conjugate gradient.
CGDs is for competitive optimization problem such as generative adversarial networks (GANs) as follows: $$ \min_{\mathbf{x}}f(\mathbf{x}, \mathbf{y}) \min_{\mathbf{y}} g(\mathbf{x}, \mathbf{y}) $$

Update: ACGD now supports distributed training. Set backward_mode=True to enable. We have new member GMRES-ACGD that can work for general two-player competitive optimization problems.

Installation

CGDs can be installed with the following pip command. It requires Python 3.6+.

pip3 install CGDs

You can also directly download the CGDs directory and copy it to your project.

Package description

The CGDs package implements the following optimization algorithms with Pytorch:

How to use

Quickstart with notebook: Examples of using ACGD.

Similar to Pytorch package torch.optim, using optimizers in CGDs has two main steps: construction and update steps.

Construction

To construct an optimizer, you have to give it two iterables containing the parameters (all should be Variables). Then you need to specify the device, learning rates.

Example:

from src import CGDs
import torch
device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
optimizer = CGDs.ACGD(max_param=model_G.parameters(), min_params=model_D.parameters(), 
                      lr_max=1e-3, lr_min=1e-3, device=device)
optimizer = CGDs.BCGD(max_params=[var1, var2], min_params=[var3, var4, var5], 
                      lr_max=0.01, lr_min=0.01, device=device)   

Update step

Both two optimizers have step() method, which updates the parameters according to their update rules. The function can be called once the computation graph is created. You have to pass in the loss but do not have to compute gradients before step() , which is different from torch.optim.

Example:

for data in dataset:
    optimizer.zero_grad()
    real_output = model_D(data)
   	latent = torch.randn((batch_size, latent_dim), device=device)
    fake_output = D(G(latent))
    loss = loss_fn(real_output, fake_output)
    optimizer.step(loss=loss)

For general competitive optimization, two losses should be defined and passed to optimizer.step

loss_x = loss_f(x, y)
loss_y = loss_g(x, y)
optimizer.step(loss_x, loss_y)

Citation

Please cite it if you find this code useful.

@misc{cgds-package,
  author = {Hongkai Zheng},
  title = {CGDs},
  year = {2020},
  publisher = {GitHub},
  journal = {GitHub repository},
  howpublished = {\url{https://github.com/devzhk/cgds-package}},
}

Metadata

Release files for CGDs 0.4.5

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for CGDs 0.4.5
File Size Uploaded
CGDs-0.4.5.tar.gz 12.8 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for CGDs 0.4.5
File Interpreter ABI Platform
CGDs-0.4.5-py3-none-any.whl Python 3 none any Details

Total release size: 27.6 kB

Release files / CGDs-0.4.5.tar.gz

Download URL CGDs-0.4.5.tar.gz
Size 12.8 kB
Tags Source
SHA-256 checksum
How to use checksums
d52db1cb346f71b887ec83ab14d478b878031381c92246f621c34f49cd907232
BLAKE2b-256 checksum
How to use checksums
6dfa205d7e1cc4753513b68958853e656721e05539713001578541dafc82379f
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/3.8.0 pkginfo/1.8.2 readme-renderer/26.0 requests/2.27.1 requests-toolbelt/0.9.1 urllib3/1.26.8 tqdm/4.46.1 importlib-metadata/4.5.0 keyring/21.2.1 rfc3986/1.4.0 colorama/0.4.3 CPython/3.7.5

Release files / CGDs-0.4.5-py3-none-any.whl

Download URL CGDs-0.4.5-py3-none-any.whl
Size 14.8 kB
Tags Python 3
SHA-256 checksum
How to use checksums
94ca09a691fe05611be2816c6cd0a76e6a508e0dac52cde7fcd950fe36442644
BLAKE2b-256 checksum
How to use checksums
8136d14c27dcf0b87e05f5061bcb868ebeda463b0d0913ea34061420c7e9333a
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/3.8.0 pkginfo/1.8.2 readme-renderer/26.0 requests/2.27.1 requests-toolbelt/0.9.1 urllib3/1.26.8 tqdm/4.46.1 importlib-metadata/4.5.0 keyring/21.2.1 rfc3986/1.4.0 colorama/0.4.3 CPython/3.7.5

Release history Release notifications | RSS feed

This release

0.4.5 This release

2 release files

0.4.4

2 release files

0.4.3

2 release files

0.4.2

2 release files

0.4.1

2 release files

0.3.0

2 release files

0.2.0

2 release files

0.1.1

2 release files

0.1.0

2 release files

0.0.4

2 release files

0.0.3

2 release files

0.0.2

2 release files

0.0.1

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page