Skip to main content

A simple implementation of control variates operator (CVor)

Project description

CVor

CVor is a Python package designed for advanced loss computation in neural networks, particularly with PyTorch. It offers a unique approach to calculating loss using three different methods: AO, NN, and LOO. CVor provides low-variance, unbiased gradient estimation based on control variates. It uniquely focuses on transforming the gradient mapping process and demonstrating superior performance in various machine learning benchmarks, including variational autoencoder training and reinforcement learning tasks.

Features

  • Multiple Loss Computation Methods: 'AO' (Average Optimization), 'NN' (Neural Network based), and 'LOO' (Leave-One-Out) methods.
  • Flexible Alpha Parameter: Allows setting the alpha parameter within the range [0, 1] for CVor loss adjustment.

Installation

You can install CVor using pip:

pip install cvor

Usage

Here's a quick example of how to use CVor:

import torch
from CVor import CVor_loss_PyTorch

# Sample loss tensor
loss_input = torch.tensor([1.0, 2.0, 3.0], device='cuda')

# Compute CVor loss
loss = CVor_loss_PyTorch(loss_input, method='NN', alpha=0.1)
print(loss)

Parameters

  • loss_input (Tensor): The input tensor representing loss.
  • method (str): The method for computing CVor loss. Options are 'AO', 'NN', and 'LOO'.
  • alpha (float): The adjustment coefficient. Must be between 0 and 1.

F_value Calculation

CVor now supports calculation of the F_value, which is a key component in computing the CVor loss. This feature allows users to inject their own logic into the loss calculation process, offering greater flexibility and adaptability to specific needs.

Using F_calculator

To use this feature, define a function that takes loss_input and alpha as parameters and returns the calculated F_value. This function can then be passed to CVor_loss_PyTorch as the F_calculator argument.

Example

Here's an example of how to define and use a F_value calculation function:

import torch
from CVor import CVor_loss_PyTorch

# F_value calculation function
def F_calculator(loss_input, alpha):
    mean_value = loss_input.mean()
    F_value = alpha * mean_value / loss_input.sum()
    return F_value

# Sample loss tensor
loss_input = torch.tensor([1.0, 2.0, 3.0], device='cuda')

# Compute CVor loss using F_value calculator
loss = CVor_loss_PyTorch(loss_input, F_calculator=F_calculator)
print(loss)

In this example, F_calculator computes F_value based on the mean of the input tensor and the alpha value. You can define the logic of F_value calculation as per your requirements.

Note

When using F_calculator, the method parameter in CVor_loss_PyTorch is ignored, and the provided function is used instead for the F_value calculation.

Requirements

  • Python 3.8+
  • PyTorch

Contributing

Contributions, issues, and feature requests are welcome. Feel free to check issues page if you want to contribute.

License

Distributed under the MIT License. See LICENSE for more information.

Contact

CHEN XINGYAN - xychen@swufe.edu.cn

Project Link: https://github.com/uglyghost/cvor.git

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

CVor-0.0.7.tar.gz (4.3 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

CVor-0.0.7-py3-none-any.whl (4.4 kB view details)

Uploaded Python 3

File details

Details for the file CVor-0.0.7.tar.gz.

File metadata

  • Download URL: CVor-0.0.7.tar.gz
  • Upload date:
  • Size: 4.3 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/4.0.2 CPython/3.10.11

File hashes

Hashes for CVor-0.0.7.tar.gz
Algorithm Hash digest
SHA256 7f1710d48f3c75099c5cc3d812089d4ac778564d1c7d0aac050b169a2deca55b
MD5 1de97442b14729dd25d2b22a0cb6ac2f
BLAKE2b-256 a9664a24b8b891142379c99aa01e14219ee86d14c5d863494e3296387a77d6c9

See more details on using hashes here.

File details

Details for the file CVor-0.0.7-py3-none-any.whl.

File metadata

  • Download URL: CVor-0.0.7-py3-none-any.whl
  • Upload date:
  • Size: 4.4 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/4.0.2 CPython/3.10.11

File hashes

Hashes for CVor-0.0.7-py3-none-any.whl
Algorithm Hash digest
SHA256 ef81f766f9a4a3247d3b96d003c9d21f1c13a8311e90d4be09c63f02de447d0d
MD5 00f4e19363e23a3cd60752b9e6518b5e
BLAKE2b-256 0be31c18f3e33142c7dd64d89fd0696091f32fe32712a63b2c51577c95e32ab4

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page