OSCR

A multi-view TSK fuzzy system with deformable Gaussian membership functions and rule-level attention for classification.

Code ↔ Paper

5 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 5 matches
  1. [1] § 4. Experimental design and results analysis › 4.2 Experimental settings ↔ pytsk_MVC/gradient_descent/training.py, lines 78–142 · score 0.76 · weight decay, AdamW, EarlyStopping, split, patience, preprocessing
  2. [2] § 3. Methodology ↔ pytsk_MVC/gradient_descent/training.py, lines 78–142 · score 0.68 · cross entropy loss, AdamW, optimizer, training, classification, model
  3. [3] § 4. Experimental design and results analysis › 4.1 Dataset overview ↔ loadDatasets.py, lines 27–65 · score 0.65 · EEG DWT WAV, Handwritten Numerals, Forest, Dermatology, Caltech7
  4. [4] § 4. Experimental design and results analysis › 4.1 Dataset overview ↔ data.py, lines 67–211 · score 0.62 · EEG DWT WAV, pipeline, Max, Min, Forest, Dermatology
  5. [5] § 4. Experimental design and results analysis › 4.2 Experimental settings ↔ data.py, lines 67–211 · score 0.58 · EarlyStopping, stratified, split, patience, PyTorch, fitted

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 558 lines · 25 KB · no license · 2 matches

  1. import torch
  2. import torch.nn as nn
  3. import torch.optim as optim
  4. import numpy as np
  5. from scipy.special import softmax
  6. from torch.utils.data import DataLoader
  7. from .utils import NumpyDataLoader
  8. # loss function
  9. #--------------------------------------tsk中地UR loss---------------------------------------------------------------------------
  10. def ur_loss(frs, tau=0.5):
  11. """
  12. The uniform regularization (UR) proposed by Cui et al. [3].
  13. UR loss is computed as :math:`\ell_{UR} = \sum_{r=1}^R (\frac{1}{N}\sum_{n=1}^N f_{n,r} - \tau)^2`,
  14. where :math:`f_{n,r}` represents the firing level of the :math:`n`-th sample on the :math:`r`-th rule.
  15. :param torch.tensor frs: The firing levels (output of the antecedent) with the size of :math:`[N, R]`,
  16. where :math:`N` is the number of samples, :math:`R` is the number of ruels.
  17. :param float tau: The expectation :math:`\tau` of the average firing level for each rule. For a
  18. :math:`C`-class classification problem, we recommend setting :math:`\tau` to :math:`1/C`,
  19. for a regression problem, :math:`\tau` can be set as :math:`0.5`.
  20. :return: A scale value, representing the UR loss.
  21. """
  22. return ((torch.mean(frs, dim=0) - tau) ** 2).sum()
  23. ## ------------------------------------------TMC中的loss function----------------------------------
  24. def KL(alpha, c):
  25. beta = torch.ones((1, c)).cuda()#创建1-tensor,用gpu计算
  26. S_alpha = torch.sum(alpha, dim=1, keepdim=True)#在某维度求和,并保持整体维度不变
  27. S_beta = torch.sum(beta, dim=1, keepdim=True)
  28. '''
  29. torch.lgamma:
  30. 一种计算公式
  31. https://pytorch.org/docs/stable/generated/torch.lgamma.html
  32. '''
  33. lnB = torch.lgamma(S_alpha) - torch.sum(torch.lgamma(alpha), dim=1, keepdim=True)
  34. lnB_uni = torch.sum(torch.lgamma(beta), dim=1, keepdim=True) - torch.lgamma(S_beta)
  35. '''
  36. torch.digamma:
  37. 计算输入的 gamma 函数的对数的导数。
  38. https://pytorch.org/docs/stable/special.html#torch.special.digamma
  39. '''
  40. dg0 = torch.digamma(S_alpha)
  41. dg1 = torch.digamma(alpha)
  42. kl = torch.sum((alpha - beta) * (dg1 - dg0), dim=1, keepdim=True) + lnB + lnB_uni
  43. return kl
  44. def ce_loss(p, alpha, c, global_step, annealing_step):#交叉熵
  45. S = torch.sum(alpha, dim=1, keepdim=True)#对输入的tensor数据的某一维度求和,keepdim=True保持求和后维度不变
  46. E = alpha - 1
  47. '''
  48. F.one_hot:
  49. 独热编码
  50. https://blog.csdn.net/qq_43760191/article/details/121778553
  51. ----------------------------------------------------------------------------------
  52. torch.digamma:
  53. 计算输入的 gamma 函数的对数的导数。
  54. https://pytorch.org/docs/stable/special.html#torch.special.digamma
  55. '''
  56. label = torch.nn.functional.one_hot(p, num_classes=c)
  57. A = torch.sum(label * (torch.digamma(S) - torch.digamma(alpha)), dim=1, keepdim=True)
  58. annealing_coef = min(1, global_step / annealing_step)#取最小值
  59. alp = E * (1 - label) + 1
  60. B = annealing_coef * KL(alp, c)
  61. return (A + B)
  62. class Wrapper:
  63. """
  64. This class provide a training framework for beginners to train their fuzzy neural networks.
  65. param torch.nn.Module model: The pre-defined TSK model.
  66. param torch.Optimizer optimizer: Pytorch optimizer.
  67. :param torch.nn._Loss: Pytorch loss. For example, :code:`torch.nn.CrossEntropyLoss()` for classification tasks,
  68. and :code:`torch.nn.MSELoss()` for regression tasks.
  69. :param int batch_size: Batch size during training & prediction.分支大小
  70. :param int epochs: Training epochs.训练阶段
  71. :param [Callback] callbacks: List of callbacks.回归列表
  72. :param str label_type: Label type, "c" or "r", when :code:`label_type="c"`, label's dtype will be changed to
  73. "int64", when :code:`label_type="r"`, label's dtype will be changed to "float32".
  74. Examples
  75. --------
  76. >>> from pytsk.gradient_descent import antecedent_init_center, AntecedentGMF, TSK, EarlyStoppingACC, EvaluateAcc, Wrapper
  77. >>> from sklearn.model_selection import train_test_split
  78. >>> from sklearn.metrics import accuracy_score
  79. >>> from sklearn.datasets import make_classification
  80. >>> from sklearn.preprocessing import StandardScaler
  81. >>> from torch.optim import AdamW
  82. >>> import torch.nn as nn
  83. >>> # ----------------- define data -----------------
  84. >>> X, y = make_classification(random_state=0)
  85. >>> x_train, x_test, y_train, y_test = train_test_split(X, y, test_size=0.3)
  86. >>> ss = StandardScaler()
  87. >>> x_train = ss.fit_transform(x_train)
  88. >>> x_test = ss.transform(x_test)
  89. >>> # ----------------- define TSK model -----------------
  90. >>> n_rule = 10 # define number of rules
  91. >>> n_class = 2 # define output dimension
  92. >>> order = 1 # first-order TSK is used
  93. >>> consbn = True # consbn tech is used
  94. >>> weight_decay = 1e-8 # weight decay for pytorch optimizer
  95. >>> lr = 0.01 # learning rate for pytorch optimizer
  96. >>> init_center = antecedent_init_center(x_train, y_train, n_rule=n_rule) # obtain the initial antecedent center
  97. >>> gmf = AntecedentGMF(in_dim=x_train.shape[1], n_rule=n_rule, high_dim=True, init_center=init_center) # define antecedent
  98. >>> model = TSK(in_dim=x_train.shape[1], out_dim=n_class, n_rule=n_rule, antecedent=gmf, order=order, consbn=consbn) # define TSK
  99. >>> # ----------------- define optimizers -----------------
  100. >>> ante_param, other_param = [], []
  101. >>> for n, p in model.named_parameters():
  102. >>> if "center" in n or "sigma" in n:
  103. >>> ante_param.append(p)
  104. >>> else:
  105. >>> other_param.append(p)
  106. >>> optimizer = AdamW(
  107. >>> [{'params': ante_param, "weight_decay": 0}, # antecedent parameters usually don't need weight_decay
  108. >>> {'params': other_param, "weight_decay": weight_decay},],
  109. >>> lr=lr
  110. >>> )
  111. >>> # ----------------- split 20% data for earlystopping -----------------
  112. >>> x_train, x_val, y_train, y_val = train_test_split(x_train, y_train, test_size=0.2)
  113. >>> # ----------------- define the earlystopping callback -----------------
  114. >>> EACC = EarlyStoppingACC(x_val, y_val, verbose=1, patience=40, save_path="tmp.pkl") # Earlystopping
  115. >>> TACC = EvaluateAcc(x_test, y_test, verbose=1) # Check test acc during training
  116. >>> # ----------------- train model -----------------
  117. >>> wrapper = Wrapper(model, optimizer=optimizer, criterion=nn.CrossEntropyLoss(),
  118. >>> epochs=300, callbacks=[EACC, TACC], ur=0, ur_tau=1/n_class) # define training wrapper, ur weight is set to 0
  119. >>> wrapper.fit(x_train, y_train) # fit
  120. >>> wrapper.load("tmp.pkl") # load best model saved by EarlyStoppingACC callback
  121. >>> y_pred = wrapper.predict(x_test).argmax(axis=1) # predict, argmax for extracting classfication label
  122. >>> print("[TSK] ACC: {:.4f}".format(accuracy_score(y_test, y_pred))) # print ACC
  123. """
  124. def __init__(self, model, optimizer, criterion,
  125. batch_size, epochs=1, callbacks=None, label_type="c",
  126. device="cuda:0", reset_param=True, ur=0, ur_tau=0.5, **kwargs):
  127. self.model = model
  128. self.optimizer = optimizer
  129. self.criterion = criterion
  130. self.batch_size = batch_size
  131. self.epochs = epochs
  132. self.device = device#设置使用gpu or cpu
  133. self.model.to(self.device)
  134. self.label_type = label_type
  135. self.ur = ur
  136. self.ur_tau = ur_tau
  137. if callbacks is None:
  138. self.callbacks = []
  139. elif isinstance(callbacks, list):#isinstance() 函数来判断一个对象是否是一个已知的类型
  140. self.callbacks = callbacks
  141. else:
  142. raise ValueError("callback must be a Callback object")
  143. self.reset_param = reset_param
  144. if self.reset_param:
  145. self.model.reset_parameters()#重新设置参数
  146. self.cur_batch = 0
  147. self.cur_epoch = 0
  148. self.kwargs = kwargs
  149. def DS_Combin(self, alpha):
  150. """
  151. :param alpha: All Dirichlet distribution parameters.
  152. :return: Combined Dirichlet distribution parameters.
  153. """
  154. def DS_Combin_two(alpha1, alpha2):
  155. """
  156. :param alpha1: Dirichlet distribution parameters of view 1
  157. :param alpha2: Dirichlet distribution parameters of view 2
  158. :return: Combined Dirichlet distribution parameters
  159. """
  160. alpha = dict()#字典
  161. alpha[0], alpha[1] = alpha1, alpha2#赋值
  162. b, S, E, u = dict(), dict(), dict(), dict()#字典
  163. for v in range(2):#
  164. S[v] = torch.sum(alpha[v], dim=1, keepdim=True)#求和,将维度1求和,并保持维度不变
  165. E[v] = alpha[v]-1#将所有值减一
  166. b[v] = E[v]/(S[v].expand(E[v].shape))#拓展维度,将维度中1拓展为E[v].shape
  167. u[v] = self.model.classes/S[v]#形成(200,1)
  168. '''
  169. torch.bmm:--矩阵相乘
  170. 要求:input 和 mat2 必须是 3-D 张量,每个张量都包含相同数量的矩阵。
  171. 如果input的维度是( b × n × m ) (b×n×m),mat2维度是( b × m × p ) (b×m×p),
  172. 那么返回的结果out就是:( b × n × p ) (b×n×p)
  173. '''
  174. # b^0 @ b^(0+1)
  175. bb = torch.bmm(b[0].view(-1, self.model.classes, 1), b[1].view(-1, 1, self.model.classes))
  176. '''
  177. torch.mul:
  178. 是矩阵a和b对应位相乘,比如a的维度是(1, 2),b的维度是(1, 2),返回的仍是(1, 2)的矩阵
  179. 区别:torch.mm(a, b)是矩阵a和b矩阵相乘,比如a的维度是(1, 2),b的维度是(2, 3),返回的就是(1, 3)的矩阵
  180. '''
  181. # b^0 * u^1
  182. uv1_expand = u[1].expand(b[0].shape)#拓展维度值为1
  183. bu = torch.mul(b[0], uv1_expand)
  184. # b^1 * u^0
  185. uv_expand = u[0].expand(b[0].shape)#拓展维度值为一
  186. ub = torch.mul(b[1], uv_expand)#矩阵a和b对应位相乘,非矩阵相乘
  187. '''
  188. torch.diagonal()
  189. 对于二维张量就是取对角线元素
  190. 对于三维张量,比如(6,m, n)
  191. torch.diagonal(tensor, dim1=-2, dim2=-1) 代表分别取6个m*n张量的对角线元素。
  192. '''
  193. # calculate C
  194. bb_sum = torch.sum(bb, dim=(1, 2), out=None)#将维度求和(200,10,10)-》(200,)
  195. bb_diag = torch.diagonal(bb, dim1=-2, dim2=-1).sum(-1)#取三维张量的对角线元素,并求和(200,10)-》(200,)
  196. C = bb_sum - bb_diag
  197. # calculate b^a
  198. b_a = (torch.mul(b[0], b[1]) + bu + ub)/((1-C).view(-1, 1).expand(b[0].shape))#torch.mul对应位相乘
  199. # calculate u^a
  200. u_a = torch.mul(u[0], u[1])/((1-C).view(-1, 1).expand(u[0].shape))
  201. # calculate new S
  202. S_a = self.model.classes / u_a
  203. # calculate new e_k
  204. e_a = torch.mul(b_a, S_a.expand(b_a.shape))
  205. alpha_a = e_a + 1
  206. return alpha_a
  207. if(len(alpha)==1):
  208. return alpha[0]
  209. for v in range(len(alpha)-1):#
  210. if v==0:
  211. alpha_a = DS_Combin_two(alpha[0], alpha[1])
  212. else:
  213. alpha_a = DS_Combin_two(alpha_a, alpha[v+1])
  214. return alpha_a
  215. def train_on_batch(self, input, target,global_step):
  216. """
  217. Define how to update a model with one batch of data.
  218. This method can be overwrite for custom training strategy.
  219. :param torch.tensor input: Feature matrix with the size of :math:`[N,D]`,
  220. :math:`N` is the number of samples, :math:`D` is the input dimension.
  221. :param torch.tensor target: Target matrix with the size of :math:`[N,C]`,
  222. :math:`C` is the output dimension.
  223. """
  224. # update model once
  225. target=target.to(self.device)
  226. for i in range(len(input)):
  227. input[i] = input[i].to(self.device)#设置训练设备
  228. #--------------------------------model训练--------------------------------------------------------
  229. outputs, frs = self.model(input, get_frs=True)
  230. #--------------------------------loss损失函数-----------------------------------------------
  231. loss = 0
  232. alpha = dict()
  233. ur_loss_value=dict()
  234. self.cur_epoch=global_step
  235. for v_num in range(len(input)):
  236. alpha[v_num] = outputs[v_num] + 1
  237. loss += ce_loss(target, alpha[v_num], self.model.out_dim, global_step, self.model.lambda_epochs)#损失函数
  238. # ----------------------------------TSK的UR-----------------------------------
  239. # if self.ur > 0:
  240. # ur_loss_value[v_num] = ur_loss(frs[v_num], self.ur_tau)
  241. # loss += self.criterion(outputs[v_num], target) + self.ur * ur_loss_value[v_num]
  242. # else:
  243. # loss += self.criterion(outputs[v_num], target)
  244. alpha_a = self.DS_Combin(alpha) # 进行了复杂的计算
  245. evidence_a = alpha_a - 1
  246. loss += ce_loss(target, alpha_a, self.model.classes, global_step, self.model.lambda_epochs)
  247. loss = torch.mean(loss) # 求平均
  248. # if self.ur > 0:
  249. # ur_loss_value = ur_loss(frs, self.ur_tau)
  250. # loss = self.criterion(outputs, target) + self.ur * ur_loss_value
  251. # else:
  252. # loss = self.criterion(outputs, target)
  253. '''
  254. optimizer.zero_grad()
  255. 意思是把梯度置零,把loss关于weight的导数变成0.
  256. 当网络参量进行反馈时,梯度是被积累的而不是被替换掉;
  257. 但是在每一个batch时毫无疑问并不需要将两个batch的梯度混合起来累积,
  258. 因此这里就需要每个batch设置一遍zero_grad 了
  259. '''
  260. self.optimizer.zero_grad()
  261. loss.backward()## 反向传播,计算当前梯度;
  262. self.optimizer.step()##根据梯度更新网络参数
  263. _, preds = torch.max(alpha_a, 1)
  264. acc = (preds == target).float().mean().item()
  265. return loss.item(), acc
  266. preds_list = []
  267. for v_num in range(len(outputs)):
  268. out = outputs[v_num] # Tensor [batch_size, n_class]
  269. _, pred = torch.max(out, 1)
  270. preds_list.append(pred)
  271. # alpha_a 已经是融合后的输出,Tensor [batch_size, n_class]
  272. _, preds = torch.max(alpha_a, 1)
  273. acc = (preds == target).float().mean().item()
  274. # if not hasattr(self, "epoch_loss"):
  275. # self.epoch_loss = 0
  276. # self.epoch_acc = 0
  277. # self.batch_count = 0
  278. def fit(self, X, y):
  279. """
  280. Train the :code:`model` with numpy array.
  281. :param numpy.array X: Feature matrix :math:`X` with the size of :math:`[N, D]`.
  282. :param numpy.array y: Label matrix :math:`Y` with the size of :math:`[N, C]`,
  283. for classification task, :math:`C=1`, for regression task, :math:`C` is the
  284. number of the output dimension of :code:`model`.
  285. """
  286. for i in range(len(X)):
  287. X[i] = X[i].astype("float32")
  288. if self.label_type == "c":
  289. y = y.astype("int64")
  290. elif self.label_type == "r":
  291. y = y.astype("float32")
  292. else:
  293. raise ValueError("label_type can only be \"c\" or \"r\"!")
  294. '''
  295. DataLoader
  296. 就是数据加载器,结合了数据集和取样器,并且可以提供多个线程处理数据集。
  297. 在训练模型时使用到此函数,用来把训练数据分成多个小组,此函数每次抛出一组数据。直至把所有的数据都抛出。就是做一个数据的初始化。
  298. batch_size:分支数(分成的小组个数)
  299. shuffle=True:表示在每个epoch重新打乱洗牌
  300. num_workers (python:int, optional) – 是否多进程读取数据(默认为0);
  301. drop_last (bool, optional) – 当样本数不能被batchsize整除时,最后一批数据是否舍弃(default: False)
  302. '''
  303. train_loader = DataLoader(
  304. NumpyDataLoader(X, y),#将 numpy 数组转换为数据加载器
  305. batch_size=self.batch_size,
  306. shuffle=self.kwargs.get("shuffle", True),
  307. num_workers=self.kwargs.get("num_workers", 0),
  308. drop_last=self.kwargs.get("drop_last", True if self.batch_size < X[0].shape[0] else False),
  309. )
  310. self.fit_loader(train_loader)
  311. return self
  312. def fit_loader(self, train_loader):
  313. """
  314. Train the :code:`model` with user-defined pytorch dataloader.使用用户定义的 pytorch 数据加载器训练。
  315. :param torch.utils.data.DataLoader train_loader: Data loader, the
  316. output of the loader should be corresponding to the inputs of :func:`train_on_batch <train_on_batch>`.
  317. For example, if dataloader has two output, then :func:`train_on_batch <train_on_batch>`
  318. should also have two inputs.
  319. """
  320. self.stop_training = False
  321. for e in range(self.epochs):
  322. self.cur_epoch = e#当前迭代次数
  323. self.epoch_loss = 0
  324. self.epoch_acc = 0
  325. self.batch_count = 0
  326. self.__run_callbacks__("on_epoch_begin")
  327. #-----------------------------TMC训练部分------------------------------------
  328. '''
  329. model.train()
  330. 就告诉了 BN 层,对之后输入的每个 batch 独立计算其均值和方差,BN 层的参数是在不断变化的。
  331. 告诉 Dropout 层,你下面应该遮住一神经元
  332. -------------------------------------------------------------------------------
  333. '''
  334. # self.model.train()
  335. # loss_meter = AverageMeter() # Computes and stores the average and current value
  336. #-------------------------------------------------------------
  337. for batch_idx, inputs in enumerate(train_loader):
  338. self.__run_callbacks__("on_batch_begin") #
  339. '''
  340. model.train()
  341. 就告诉了 BN 层,对之后输入的每个 batch 独立计算其均值和方差,BN 层的参数是在不断变化的。
  342. 告诉 Dropout 层,你下面应该遮住一神经元
  343. -------------------------------------------------------------------------------
  344. '''
  345. self.model.train()
  346. # self.train_on_batch(inputs[0], inputs[1], e) # 自定义的训练策略
  347. batch_loss, batch_acc = self.train_on_batch(inputs[0], inputs[1],self.cur_epoch) # 自定义的训练策略
  348. self.epoch_loss += batch_loss
  349. self.epoch_acc += batch_acc
  350. self.batch_count += 1
  351. self.__run_callbacks__("on_batch_end")
  352. self.cur_batch += 1
  353. self.epoch_loss /= self.batch_count
  354. self.epoch_acc /= self.batch_count
  355. self.__run_callbacks__("on_epoch_end")
  356. if self.stop_training:
  357. break
  358. return self
  359. def predict(self, X,y,global_step):
  360. """
  361. Get the prediction of the model.
  362. :param global_step:
  363. :param numpy.array X: Feature matrix :math:`X` with the size of :math:`[N, D]`.
  364. :param y: Not used.
  365. :return: Prediction matrix :math:`\hat{Y}` with the size of :math:`[N, C]`,
  366. :math:`C` is the output dimension of the :code:`model`.输出预测的标签y
  367. """
  368. for i in range(len(X)):
  369. X[i] = X[i].astype("float32")
  370. if self.label_type == "c":
  371. y = y.astype("int64")
  372. elif self.label_type == "r":
  373. y = y.astype("float32")
  374. else:
  375. raise ValueError("label_type can only be \"c\" or \"r\"!")
  376. test_loader = DataLoader(
  377. NumpyDataLoader(X,y),
  378. batch_size=self.batch_size,#分支数(分成的小组个数)
  379. shuffle=False,#表示在每个epoch不重新打乱洗牌
  380. num_workers=self.kwargs.get("num_workers", 0),# 是否多进程读取数据(默认为0)
  381. drop_last=False#当样本数不能被batchsize整除时,最后一批数据是否舍弃(default: False)
  382. )
  383. Class_preds =[]#用来村类别预测
  384. loss_meter = AverageMeter()#用来计算loss的平均值
  385. #----【遍历数据】:将数据分为X个200小组,遍历X次-----------------------------------
  386. for batch_idx, (inputs,target) in enumerate(test_loader):
  387. self.model.eval()#告诉BN层不要再变,Dropout层别遮挡当神经元
  388. for i in range(len(inputs)):
  389. inputs[i] = inputs[i].to(self.device) # inputs[0]:字典---设置训练设备
  390. # --------------------------------model预测--------------------------------------------------------
  391. outputs = self.model(inputs)#inputs[0]:字典--多视图数据
  392. # --------------------------------loss损失函数-----------------------------------------------
  393. loss = 0
  394. alpha = dict()
  395. target = target.to(self.device)
  396. for v_num in range(len(inputs)):
  397. alpha[v_num] = outputs[v_num] + 1
  398. loss += ce_loss(target, alpha[v_num], self.model.out_dim, global_step, self.model.lambda_epochs) # 损失函数
  399. alpha_a = self.DS_Combin(alpha) # 进行了复杂的计算--从6个视图结合--》》变为一个视图数据
  400. evidence_a = alpha_a - 1
  401. loss += ce_loss(target, alpha_a, self.model.classes, global_step, self.model.lambda_epochs)
  402. loss = torch.mean(loss) # 求平均
  403. #----------------挑【evidence_a.data】中【维度为1】最大值作为类别的确定条件-----------------------------------
  404. _, predicted = torch.max(evidence_a.data, 1) # 返回输入张量维度为1中的元素最大值。
  405. #------------------------------TMC 计算正确数量,本函数就能计算正确率--------------------------------------
  406. # correct_num += (predicted == target).sum().item() # 记录正确的次数
  407. loss_meter.update(loss.item())
  408. #--------------------------TSk 将预测结果存起来--给上级函数计算正确率----------------------------------------------
  409. '''
  410. detach().cpu().numpy():
  411. 阻断反向传播(因为这里是测试)-将数据从GPU转向CPU-tensor变量转numpy
  412. 解释:https://blog.csdn.net/weixin_38424903/article/details/107649436
  413. '''
  414. predicted =predicted.detach().cpu().numpy()
  415. Class_preds.append(predicted)#将小组(每个小组200个样本数)的训练后数据存进对应列表(对应视图)中
  416. '''
  417. -----结束循环后,获得X个分支预测类别的list,将list合并成一个总和预测数据----------------------------------------
  418. numpy.concatenate((a1,a2,...), axis=0)
  419. 函数。能够一次完成多个数组的拼接
  420. https://www.cnblogs.com/shueixue/p/10953699.html
  421. '''
  422. return np.concatenate(Class_preds, axis=0),loss_meter.avg
  423. def predict_proba(self, X, y=None):
  424. """
  425. For classification problem only, need :code:`label_type="c"`, return the prediction after softmax.
  426. :param numpy.array X: Feature matrix :math:`X` with the size of :math:`[N, D]`.
  427. :param y: Not used.
  428. :return: Prediction matrix :math:`\hat{Y}` with the size of :math:`[N, C]`,
  429. :math:`C` is the output dimension of the :code:`model`.
  430. """
  431. if self.label_type == "r":
  432. raise ValueError("predict_proba can only be used when label_type=\"c\"")
  433. y_preds = self.predict(X)
  434. return softmax(y_preds, axis=1)
  435. def __run_callbacks__(self, func_name):
  436. for cb in self.callbacks:
  437. getattr(cb, func_name)(self)#getattr() 函数用于返回一个对象属性值
  438. def save(self, path):
  439. """
  440. Save model.
  441. :param str path: Model save path.
  442. """
  443. torch.save(self.model.state_dict(), path)
  444. def load(self, path):
  445. """
  446. Load model.
  447. :param str path: Model save path.
  448. """
  449. self.model.load_state_dict(torch.load(path))
  450. class AverageMeter(object):
  451. """Computes and stores the average and current value"""
  452. def __init__(self):
  453. self.reset()
  454. def reset(self):
  455. self.val = 0
  456. self.avg = 0
  457. self.sum = 0
  458. self.count = 0
  459. def update(self, val, n=1):
  460. self.val = val
  461. self.sum += val * n
  462. self.count += n
  463. self.avg = self.sum / self.count

training.py at commit d466603, no license · at the source

Overview

Authors: Zhiqi Huang1,2, Yizhang Jiang1,2, Kaijian Xia2,3,4
ORCID iDs: Zhiqi Huang
  1. School of Artificial Intelligence and Computer Science, Jiangnan University, Wuxi, Jiangsu, China
  2. Engineering Research Center of the Ministry of Education for Intelligent Technology and Healthcare, Jiangnan University, P.R. China
  3. Changshu Key Laboratory of Medical Affiliated Intelligence and Big Data, Suzhou, Jiangsu, China
  4. Center of Intelligent Medical Technology Research, Changshu Hospital Affiliated to Soochow University, Suzhou, Jiangsu, China
Institutions: Jiangnan University (China); Soochow University (China)
Journal: PloS one, volume 21, issue 5, article e0348610
Dates: received 9 September 2025; accepted 17 April 2026; published online 11 May 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pone.0348610 · PMID 42113778 · PMCID PMC13160309 · OpenAlex W7160848192
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: methods / tools (subfield)
Methods: Statistics, Graphs, Machine learning
MeSH: Classification Algorithms*, Fuzzy Logic*, Normal Distribution*, Datasets as Topic (* major topic)
Topic: Fuzzy Logic and Control Systems (Artificial Intelligence, Computer Science), according to OpenAlex
Citations: not cited yet (Europe PMC); 21 references in the paper

Abstract

This study presents a novel multi-view TSK fuzzy system that integrates deformable Gaussian membership functions with a rule-level attention mechanism (MDA-TSK-FS), aiming to improve the modeling capacity and flexibility of fuzzy systems in high-dimensional and complex classification tasks. In the antecedent part, learnable deformation offsets are introduced, enabling the membership function centers of each fuzzy rule to dynamically adjust according to data characteristics. This design enhances the adaptability of rules to the input space. Furthermore, a multi-head attention mechanism at the rule level is incorporated to adaptively allocate rule weights based on sample-specific information, thereby enabling dynamic modeling of rule importance and optimized rule selection. Extensive experiments on five public multi-view datasets, including Caltech7, Handwritten, Dermatology, Forest, and EEG, demonstrate that the proposed model consistently achieves superior performance, reaching classification accuracies of 94.38%, 98.62%, 98.58%, 88.57%, and 69.75%, respectively, and outperforming strong baselines. Ablation studies further verify the effectiveness of the two core components: the deformable antecedent structure and the rule-level attention mechanism, which individually improved EEG classification accuracy by approximately 7% and 6%, and jointly by 9.25% compared to the baseline. Notably, the model exhibits superior generalization and interpretability, particularly when processing multi-source heterogeneous data. These findings indicate that the proposed approach provides a new modeling paradigm for multi-view fuzzy inference, offering both theoretical contributions and practical application potential.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 5 matches between paragraphs and lines of code.

marisazc010/MDA-TSK-FS

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: d4666030a04d014c757833e1e47a450ad03a3303, 4 March 2026
Languages: Python (20)
Size: 60 files, 20 scripts
Software Heritage: not archived
Found in: “Data Availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: PyTorch (12 files), NumPy (11 files), scikit-learn (6 files), Matplotlib (4 files), SciPy (4 files), h5py (2 files), pandas (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
20 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 20 scripts, each with its path and the digest of its content;
  • 5 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data Availability

All datasets evaluated in this study are publicly accessible from their original sources, including the Caltech Vision Dataset, the UCI Machine Learning Repository, and the PhysioNet Database. To ensure full transparency and reproducibility, the standardized. mat data files, feature representations, and the complete preprocessing scripts are publicly available from the GitHub repository. The complete codebase and materials can be accessed from the GitHub repository (https://github.com/marisazc010/MDA-TSK-FS).

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 2, 28 September 2026

  • Funding: added National Natural Science Foundation of China: 62171203; Government of Jiangsu Province

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 3 authors, 4 MeSH terms, 11 references.

Cite

This paper

Huang, Z., Jiang, Y., & Xia, K. (2026). A multi-view TSK fuzzy system with deformable Gaussian membership functions and rule-level attention for classification. PloS one, 21(5), e0348610. https://doi.org/10.1371/journal.pone.0348610

BibTeX

@article{huang2026multi,
author = {Huang, Zhiqi and Jiang, Yizhang and Xia, Kaijian},
title = {{A multi-view TSK fuzzy system with deformable Gaussian membership functions and rule-level attention for classification}},
journal = {PloS one},
year = {2026},
month = may,
volume = {21},
number = {5},
pages = {e0348610},
publisher = {PLOS},
issn = {1932-6203},
doi = {10.1371/journal.pone.0348610},
url = {https://doi.org/10.1371/journal.pone.0348610},
pmid = {42113778},
pmcid = {PMC13160309}
}

RIS

TY - JOUR
AU - Huang, Zhiqi
AU - Jiang, Yizhang
AU - Xia, Kaijian
TI - A multi-view TSK fuzzy system with deformable Gaussian membership functions and rule-level attention for classification
T2 - PloS one
J2 - PLoS One
PY - 2026
DA - 2026/05/11
VL - 21
IS - 5
SP - e0348610
SN - 1932-6203
PB - PLOS
DO - 10.1371/journal.pone.0348610
UR - https://doi.org/10.1371/journal.pone.0348610
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pone.0348610",
"type": "article-journal",
"title": "A multi-view TSK fuzzy system with deformable Gaussian membership functions and rule-level attention for classification",
"container-title": "PloS one",
"author": [
{
"family": "Huang",
"given": "Zhiqi"
},
{
"family": "Jiang",
"given": "Yizhang"
},
{
"family": "Xia",
"given": "Kaijian"
}
],
"container-title-short": "PLoS One",
"volume": "21",
"issue": "5",
"page": "e0348610",
"DOI": "10.1371/journal.pone.0348610",
"PMID": "42113778",
"PMCID": "PMC13160309",
"ISSN": "1932-6203",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pone.0348610",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
11
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s42003-026-10169-0 [code]
Shared representations in brains and models reveal a two-route cortical organization during scene perception.
Journal: Communications biology
In common: h5py, PyTorch, scikit-learn, 4 other tools, 1 reference
[2] doi:10.7554/elife.110588 [code]
Opening the black box toward a modular approach to spike sorting.
Journal: eLife
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[3] doi:10.1002/advs.77003 [code]
SemanticST: A Scalable Multi-Contextual Graph Learning Framework for Uncovering Spatial Niches and Robust Multi-Sample Integration in Spatial Transcriptomics.
Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[4] doi:10.1038/s41467-026-76939-w [code]
HIPPIE: a generative model for electrophysiological analysis across species, technologies, and modalities.
Journal: Nature communications
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[5] doi:10.1093/bioinformatics/btag540 [code]
Deciphering spatial heterogeneity by multimodal spatial transcriptomics modelling with SpatialModal.
Journal: Bioinformatics (Oxford, England)
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[6] doi:10.1007/s12021-026-09803-3 [code]
NeuroFusion: A Unified Framework for Generalized Visual Stimulus Decoding from fMRI Across Datasets and Subjects.
Journal: Neuroinformatics
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[7] doi:10.1162/imag.a.1299 [code]
A modular semantic-structural pipeline for visual decoding from primate spiking data via selective temporal integration.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[8] doi:10.1016/j.isci.2026.116671 [code]
A high-resolution functional network-organized atlas of human superficial white matter from ultra-high-field diffusion MRI.
Journal: iScience
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[9] doi:10.1038/s41598-026-57519-w [code]
Automated segmentation of neurons and spinal cord structures in immunofluorescence images using SpineDL.
Journal: Scientific reports
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools
[10] doi:10.21203/rs.3.rs-9676637/v1 [code]
A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic Technologies
Journal: Research Square (preprint)
In common: h5py, PyTorch, scikit-learn, 4 other tools, methods / tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.