Browse Source

[MNT] Solve the conlict

tags/v0.3.2
bxdd 3 years ago
parent
commit
fd84ddddbc
1 changed files with 16 additions and 17 deletions
  1. +16
    -17
      learnware/specification/rkme.py

+ 16
- 17
learnware/specification/rkme.py View File

@@ -1,4 +1,5 @@
import os
import mkl
import copy
import torch
import faiss
@@ -11,16 +12,17 @@ from typing import Tuple, Any, List, Union, Dict

from .base import BaseStatSpecification

mkl.get_max_threads()

def setup_seed(seed):
"""
Fix a random seed for addressing reproducibility issues.

Parameters
----------
seed : int
Random seed for torch, torch.cuda, numpy, random and cudnn libraries.
"""
def setup_seed(seed):
"""Fix a random seed for addressing reproducibility issues.
Parameters
----------
seed : int
Random seed for torch, torch.cuda, numpy, random and cudnn libraries.
"""
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
np.random.seed(seed)
@@ -29,8 +31,7 @@ def setup_seed(seed):


def choose_device(cuda_idx=-1):
"""
Let users choose compuational device between CPU or GPU.
"""Let users choose compuational device between CPU or GPU.

Parameters
----------
@@ -50,8 +51,7 @@ def choose_device(cuda_idx=-1):


def torch_rbf_kernel(x1, x2, gamma) -> torch.Tensor:
"""
Use pytorch to compute rbf_kernel function at faster speed.
"""Use pytorch to compute rbf_kernel function at faster speed.

Parameters
----------
@@ -74,11 +74,10 @@ def torch_rbf_kernel(x1, x2, gamma) -> torch.Tensor:


def solve_qp(K: np.ndarray, C: np.ndarray):
"""
Solver for the following quadratic programming(QP) problem:
- min 1/2 x^T K x - C^T x
s.t 1^T x - 1 = 0
- I x <= 0
"""Solver for the following quadratic programming(QP) problem:
- min 1/2 x^T K x - C^T x
s.t 1^T x - 1 = 0
- I x <= 0

Parameters
----------


Loading…
Cancel
Save