mindspore/docs/api/api_python/probability/mindspore.nn.probability.to...

46 lines
1.7 KiB
ReStructuredText
Raw Normal View History

2022-05-27 14:23:35 +08:00
mindspore.nn.probability.toolbox.VAEAnomalyDetection
====================================================
.. py:class:: mindspore.nn.probability.toolbox.VAEAnomalyDetection(encoder, decoder, hidden_size=400, latent_size=20)
使用 VAE 进行异常检测的工具箱。
2022-07-01 10:49:56 +08:00
变分自动编码器VAE可用于无监督异常检测。异常分数是 sample_x 与重建 sample_x 之间的误差。如果分数高,则 X 大多是异常值。
2022-05-27 14:23:35 +08:00
2022-07-18 15:22:30 +08:00
参数:
- **encoder** (Cell) - 定义为编码器的深度神经网络 (DNN) 模型。
- **decoder** (Cell) - 定义为解码器的深度神经网络 (DNN) 模型。
- **hidden_size** (int) - 编码器输出 Tensor 的大小。默认值400。
- **latent_size** (int) - 潜在空间的大小。默认值20。
2022-05-27 14:23:35 +08:00
.. py:method:: predict_outlier(sample_x, threshold=100.0)
预测样本是否为异常值。
2022-07-18 15:22:30 +08:00
参数:
- **sample_x** (Tensor) - 待预测的样本shape 为 (N, C, H, W)。
- **threshold** (float) - 异常值的阈值。默认值100.0。
2022-05-27 14:23:35 +08:00
2022-07-18 15:22:30 +08:00
返回:
bool样本是否为异常值。
2022-05-27 14:23:35 +08:00
.. py:method:: predict_outlier_score(sample_x)
预测异常值分数。
2022-07-18 15:22:30 +08:00
参数:
- **sample_x** (Tensor) - 待预测的样本shape 为 (N, C, H, W)。
2022-05-27 14:23:35 +08:00
2022-07-18 15:22:30 +08:00
返回:
float样本的预测异常值分数。
2022-05-27 14:23:35 +08:00
2022-06-10 08:30:45 +08:00
.. py:method:: train(train_dataset , epochs=5)
2022-05-27 14:23:35 +08:00
训练 VAE 模型。
2022-07-18 15:22:30 +08:00
参数:
- **train_dataset** (Dataset) - 用于训练模型的数据集迭代器。
- **epochs** (int) - 数据的迭代总数。默认值5。
2022-05-27 14:23:35 +08:00
2022-07-18 15:22:30 +08:00
返回:
Cell训练完的模型。