【域适应】基于散度成分分析(SCA)的四分类任务典型方法实现

2024-04-11 21:20

本文主要是介绍【域适应】基于散度成分分析(SCA)的四分类任务典型方法实现,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

关于

SCA(scatter component analysis)是基于一种简单的几何测量,即分散,它在再现内核希尔伯特空间上进行操作。 SCA找到一种在最大化类的可分离性、最小化域之间的不匹配和最大化数据的可分离性之间进行权衡的表示;每一个都通过分散进行量化。 

参考论文:Shibboleth Authentication Request

工具

MATLAB

方法实现

SCA变换实现
function [test_accuracy, predicted_labels, Zs, Zt] = SCA(X_s_cell, Y_s_cell, X_t, Y_t, params)INPUT(params is optional):X_s_cell          - cell of (n_s*d) matrix, each matrix corresponds to the instance features of a source domainY_s_cell          - cell of (n_s*1) matrix, each matrix corresponds to the instance labels of a source domainX_t               - (n_t*d) matrix, rows correspond to instances and columns correspond to featuresY_t               - (n_t*1) matrix, each row is the class label of corresponding instances in X_t[params]          - params.beta:      vector of validated values of betaparams.delta:     vector of validated values of deltaparams.k_list:    vector of validated dimension of the transformed spaceparams.X_v:       (n_v*d) matrix of instance features of validation set (use the source instances if not provided)params.Y_v:       (n_v*1) matrix of instance labels of validation set (use the source instances if not provided)params.verbose:   if true, show the validation accuracy of each parameter settingOUTPUT:test_accuracy     - test accuracy on target instancespredicted_labels  - predicted labels of target instancesZs                - projected source domain instancesZt                - projected target domain instancesShoubo Hu (shoubo.sub [at] gmail.com)
2019-06-02Reference
[1] Ghifary, M., Balduzzi, D., Kleijn, W. B., & Zhang, M. (2017). Scatter component analysis: A unified framework for domain adaptation and domain generalization. IEEE transactions on pattern analysis and machine intelligence, 39(7), 1414-1430.
%}if nargin < 4error('Error. \nOnly %d input arguments! At least 4 required', nargin);elseif nargin == 4% default params valuesbeta = [0.1 0.3 0.5 0.7 0.9];delta = [1e-3 1e-2 1e-1 1 1e1 1e2 1e3 1e4 1e5 1e6];k_list = [2];X_v = cat(1, X_s_cell{:});Y_v = cat(1, Y_s_cell{:});verbose = false;elseif nargin == 5if ~isfield(params, 'beta')beta = [0.1 0.3 0.5 0.7 0.9];elsebeta = params.beta;endif ~isfield(params, 'delta')delta = [1e-3 1e-2 1e-1 1 1e1 1e2 1e3 1e4 1e5 1e6];elsedelta = params.delta;endif ~isfield(params, 'k_list')k_list = [2];elsek_list = params.k_list;endif ~isfield(params, 'verbose')verbose = false;elseverbose = params.verbose;endif ~isfield(params, 'X_v')X_v = cat(1, X_s_cell{:});Y_v = cat(1, Y_s_cell{:});elseif ~isfield(params, 'Y_v')error('Error. Labels of validation set needed!');endX_v = params.X_v;Y_v = params.Y_v;endend% ----- training phase% ----- ----- source domainsX_s = cat(1, X_s_cell{:});Y_s = cat(1, Y_s_cell{:});fprintf('Number of source domains: %d, Number of classes: %d.\n', length(X_s_cell), length(unique(Y_s)) );fprintf('Validating hyper-parameters ...\n');dist_s_s = pdist2(X_s, X_s);dist_s_s = dist_s_s.^2;sgm_s = compute_width(dist_s_s);% ----- ----- validation setdist_s_v = pdist2(X_s, X_v);dist_s_v = dist_s_v.^2;sgm_v = compute_width(dist_s_s);n_s = size(X_s, 1);n_v = size(X_v, 1);H_s = eye(n_s) - ones(n_s)./n_s;H_v = eye(n_v) - ones(n_v)./n_v;K_s_s = exp(-dist_s_s./(2 * sgm_s * sgm_s));K_s_v = exp(-dist_s_v./(2 * sgm_v * sgm_v));K_s_v_bar = H_s * K_s_v * H_v;[P, T, D, Q, K_s_s_bar] = SCA_terms(K_s_s, X_s_cell, Y_s_cell);acc_mat = zeros(length(k_list), length(beta), length(delta));for i = 1:length(beta)cur_beta = beta(i);for j = 1:length(delta)cur_delta = delta(j);[B, A] = SCA_trans(P, T, D, Q, K_s_s_bar, cur_beta, cur_delta, 1e-5);for k = 1:length(k_list)[acc, ~, ~, ~] = SCA_test(B, A, K_s_s_bar, K_s_v_bar, Y_s, Y_v, k_list( k ) );acc_mat(k, i, j) = acc;if verbosefprintf('beta: %f, delta: %f, acc: %f\n', cur_beta, cur_delta, acc);endendendendfprintf('Validation done! Classifying the target domain instances ...\n');% ----- test phase% ----- ----- get optimal parametersacc_tr_best = max( acc_mat(:) );ind = find( acc_mat == acc_tr_best );[k, i, j] = size( acc_mat );[best_k, best_i, best_j] = ind2sub([k, i, j], ind(1));best_beta = beta(best_i);best_delta = delta(best_j);best_k = k_list(best_k);% ----- ----- test on the target domaindist_s_t = pdist2(X_s, X_t);dist_s_t = dist_s_t.^2;sgm = compute_width(dist_s_t);K_s_t = exp(-dist_s_t./(2 * sgm * sgm));n_s = size(X_s, 1);H_s = eye(n_s) - ones(n_s)./n_s;n_t = size(X_t, 1);H_t = eye(n_t) - ones(n_t)./n_t;K_s_t_bar = H_s * K_s_t * H_t;[B, A] = SCA_trans(P, T, D, Q, K_s_s_bar, best_beta, best_delta, 1e-5);[test_accuracy, predicted_labels, Zs, Zt] = SCA_test(B, A, K_s_s_bar, K_s_t_bar, Y_s, Y_t, best_k );fprintf('Test accuracy: %f\n', test_accuracy);end
基于SCA的域迁移分类实现
clear all
clcaddpath('./modules');
load('./syn_data/data.mat');% ----- parameters
% target / all / source domains
tgt_dm = [5];
val_dm = [3 4];
src_dm = [1 2];data_cell = XY_cell;
X_t = data_cell{tgt_dm(1)}(:, 1:2);
Y_t = data_cell{tgt_dm(1)}(:, 3);% ----- training data
X_s_cell = cell(1,length(src_dm));
Y_s_cell = cell(1,length(src_dm));    
for idx = 1:length(src_dm)cu_dm = src_dm(1, idx);X_s_cell{idx} = data_cell{cu_dm}(:, 1:2);Y_s_cell{idx} = data_cell{cu_dm}(:, 3);
end
% ----- validation data
X_v = [];
Y_v = [];
for idx = 1:length(val_dm)cu_dm = val_dm(1, idx);X_v = [X_v; data_cell{cu_dm}(:, 1:2)];Y_v = [Y_v; data_cell{cu_dm}(:, 3)];
endparams.X_v = X_v;
params.Y_v = Y_v;
params.verbose = true;
[test_accuracy, predicted_labels, Zs, Zt] = SCA(X_s_cell, Y_s_cell, X_t, Y_t, params);

代码获取

相关问题和代码开发,可后台私信沟通交流。

这篇关于【域适应】基于散度成分分析(SCA)的四分类任务典型方法实现的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/895216

相关文章

pandas中位数填充空值的实现示例

《pandas中位数填充空值的实现示例》中位数填充是一种简单而有效的方法,用于填充数据集中缺失的值,本文就来介绍一下pandas中位数填充空值的实现,具有一定的参考价值,感兴趣的可以了解一下... 目录什么是中位数填充?为什么选择中位数填充?示例数据结果分析完整代码总结在数据分析和机器学习过程中,处理缺失数

Golang HashMap实现原理解析

《GolangHashMap实现原理解析》HashMap是一种基于哈希表实现的键值对存储结构,它通过哈希函数将键映射到数组的索引位置,支持高效的插入、查找和删除操作,:本文主要介绍GolangH... 目录HashMap是一种基于哈希表实现的键值对存储结构,它通过哈希函数将键映射到数组的索引位置,支持

Java学习手册之Filter和Listener使用方法

《Java学习手册之Filter和Listener使用方法》:本文主要介绍Java学习手册之Filter和Listener使用方法的相关资料,Filter是一种拦截器,可以在请求到达Servl... 目录一、Filter(过滤器)1. Filter 的工作原理2. Filter 的配置与使用二、Listen

Pandas使用AdaBoost进行分类的实现

《Pandas使用AdaBoost进行分类的实现》Pandas和AdaBoost分类算法,可以高效地进行数据预处理和分类任务,本文主要介绍了Pandas使用AdaBoost进行分类的实现,具有一定的参... 目录什么是 AdaBoost?使用 AdaBoost 的步骤安装必要的库步骤一:数据准备步骤二:模型

Pandas统计每行数据中的空值的方法示例

《Pandas统计每行数据中的空值的方法示例》处理缺失数据(NaN值)是一个非常常见的问题,本文主要介绍了Pandas统计每行数据中的空值的方法示例,具有一定的参考价值,感兴趣的可以了解一下... 目录什么是空值?为什么要统计空值?准备工作创建示例数据统计每行空值数量进一步分析www.chinasem.cn处

使用Pandas进行均值填充的实现

《使用Pandas进行均值填充的实现》缺失数据(NaN值)是一个常见的问题,我们可以通过多种方法来处理缺失数据,其中一种常用的方法是均值填充,本文主要介绍了使用Pandas进行均值填充的实现,感兴趣的... 目录什么是均值填充?为什么选择均值填充?均值填充的步骤实际代码示例总结在数据分析和处理过程中,缺失数

Java对象转换的实现方式汇总

《Java对象转换的实现方式汇总》:本文主要介绍Java对象转换的多种实现方式,本文通过实例代码给大家介绍的非常详细,对大家的学习或工作具有一定的参考借鉴价值,需要的朋友参考下吧... 目录Java对象转换的多种实现方式1. 手动映射(Manual Mapping)2. Builder模式3. 工具类辅助映

Go语言开发实现查询IP信息的MCP服务器

《Go语言开发实现查询IP信息的MCP服务器》随着MCP的快速普及和广泛应用,MCP服务器也层出不穷,本文将详细介绍如何在Go语言中使用go-mcp库来开发一个查询IP信息的MCP... 目录前言mcp-ip-geo 服务器目录结构说明查询 IP 信息功能实现工具实现工具管理查询单个 IP 信息工具的实现服

SpringBoot基于配置实现短信服务策略的动态切换

《SpringBoot基于配置实现短信服务策略的动态切换》这篇文章主要为大家详细介绍了SpringBoot在接入多个短信服务商(如阿里云、腾讯云、华为云)后,如何根据配置或环境切换使用不同的服务商,需... 目录目标功能示例配置(application.yml)配置类绑定短信发送策略接口示例:阿里云 & 腾

Windows 上如果忘记了 MySQL 密码 重置密码的两种方法

《Windows上如果忘记了MySQL密码重置密码的两种方法》:本文主要介绍Windows上如果忘记了MySQL密码重置密码的两种方法,本文通过两种方法结合实例代码给大家介绍的非常详细,感... 目录方法 1:以跳过权限验证模式启动 mysql 并重置密码方法 2:使用 my.ini 文件的临时配置在 Wi