Skip to main content

Generalized Distance Weighted Discrimination (DWD) with sGS-ADDM

Project description

genDWD: Generalized Distance Weighted Discrimination

genDWD is a high-performance Python implementation of the Generalized Distance Weighted Discrimination (DWD) algorithm.

While Support Vector Machines (SVM) are widely used, they often suffer from the "data piling" phenomenon in HDLLS (High-Dimension, Low Large Sample) settings, where many data points project onto the same location on the decision boundary, leading to poor generalization. DWD was specifically developed to overcome this by accounting for the relative distance of all data points, providing better robustness in high-dimensional spaces.

Originally formulated for binary classification, this implementation extends the algorithm's utility by featuring built-in Multiclass classification support, enabling its application to complex, real-world datasets with multiple categories. It is designed to handle these tasks efficiently using the sGS-ADMM framework for large-scale optimization.


Core Research Basis

This implementation is strictly constructed based on the numerical procedures described in the following scientific paper:

"Fast algorithms for large scale generalized distance weighted discrimination"

— Xin Yee Lam, J.S. Marron, Defeng Sun, and Kim-Chuan Toh (September 5, 2018)

Key Features

  • Automatic Detection: Automatically switches between Binary and Multiclass (One-vs-One) logic.
  • Adaptive Solvers: Dynamically selects between Cholesky, SMW, or PSQMR based on $n$ and $d$.
  • sGS-ADMM Framework: Ensures fast convergence using Symmetric Gauss-Seidel ADMM.
  • Penalty Tuning: Automatic $C$ parameter calculation using median inter-class distance.

Quick Start

  1. Install the package:
pip install genDWD

Implementation Guide

To use the genDWD class in your project, follow the instructions below. The model follows the standard .fit() and .predict() pattern.

1. Basic Initialization

You can customize the model during initialization. The parameter C='auto' is recommended as it calculates the penalty based on the dataset's geometry.

from gendwd import genDWD

model = genDWD(C='auto')

2. Model Training and Prediction

The model follows the familiar .fit() and .predict() workflow. It automatically handles both binary and multiclass data.

Training the Model

To train the model, pass your feature matrix X and label vector y. The algorithm will automatically detect if it should use a binary or multiclass (One-vs-One) strategy.

import numpy as np

# Sample Data: 100 samples, 10 features
X_train = np.random.randn(100, 10)
# Labels can be integers, strings, or floats
y_train = np.random.choice(['Class_A', 'Class_B'], size=100)

# Train the model
model.fit(X_train, y_train)

# New data for prediction
X_new = np.random.randn(5, 10)

# Get predicted class labels
predictions = model.predict(X_new)

print(f"Predicted Labels: {predictions}")

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

gendwd-0.1.0.tar.gz (7.6 kB view details)

Uploaded Source

Built Distribution

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

gendwd-0.1.0-py3-none-any.whl (7.5 kB view details)

Uploaded Python 3

File details

Details for the file gendwd-0.1.0.tar.gz.

File metadata

  • Download URL: gendwd-0.1.0.tar.gz
  • Upload date:
  • Size: 7.6 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.13.1

File hashes

Hashes for gendwd-0.1.0.tar.gz
Algorithm Hash digest
SHA256 f1e18c63ccdc2addff52c066b3c8bbf602860f5b5fea2592d3341f48519bab89
MD5 a395bfc980b9981215ce4a699cfb9e0f
BLAKE2b-256 2789671b3c7ae614d3560a97a0596b693ca4f01f92b77e865688313d6656fea9

See more details on using hashes here.

File details

Details for the file gendwd-0.1.0-py3-none-any.whl.

File metadata

  • Download URL: gendwd-0.1.0-py3-none-any.whl
  • Upload date:
  • Size: 7.5 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.13.1

File hashes

Hashes for gendwd-0.1.0-py3-none-any.whl
Algorithm Hash digest
SHA256 80e4324b2de7aaae8b395b9649f2609a64dfc75564ae87f559297c1de6db792b
MD5 7acf3cce6236f5807e671284881045e1
BLAKE2b-256 fb75752a92e777ff78bb3fd63d28f3413036d49338e8c932939292ec1e2b2304

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