use crate::metric::{
AccuracyInput, Adaptor, ConfusionStatsInput, HammingScoreInput, LossInput, PerplexityInput,
TopKAccuracyInput, processor::ItemLazy,
};
use burn_core::tensor::{Int, Tensor};
#[derive(new)]
pub struct ClassificationOutput {
pub loss: Tensor<1>,
pub output: Tensor<2>,
pub targets: Tensor<1, Int>,
}
impl ItemLazy for ClassificationOutput {
fn sync(self) -> Self {
self.loss.device().flush();
ClassificationOutput {
output: self.output.no_grad(),
loss: self.loss.no_grad(),
targets: self.targets,
}
}
}
impl Adaptor<AccuracyInput> for ClassificationOutput {
fn adapt(&self) -> AccuracyInput {
AccuracyInput::new(self.output.clone(), self.targets.clone())
}
}
impl Adaptor<LossInput> for ClassificationOutput {
fn adapt(&self) -> LossInput {
LossInput::new(self.loss.clone())
}
}
impl Adaptor<TopKAccuracyInput> for ClassificationOutput {
fn adapt(&self) -> TopKAccuracyInput {
TopKAccuracyInput::new(self.output.clone(), self.targets.clone())
}
}
impl Adaptor<PerplexityInput> for ClassificationOutput {
fn adapt(&self) -> PerplexityInput {
PerplexityInput::new(self.output.clone(), self.targets.clone())
}
}
impl Adaptor<ConfusionStatsInput> for ClassificationOutput {
fn adapt(&self) -> ConfusionStatsInput {
let [_, num_classes] = self.output.dims();
if num_classes > 1 {
ConfusionStatsInput::new(
self.output.clone(),
self.targets.clone().one_hot(num_classes).bool(),
)
} else {
ConfusionStatsInput::new(
self.output.clone(),
self.targets.clone().unsqueeze_dim(1).bool(),
)
}
}
}
#[derive(new)]
pub struct MultiLabelClassificationOutput {
pub loss: Tensor<1>,
pub output: Tensor<2>,
pub targets: Tensor<2, Int>,
}
impl ItemLazy for MultiLabelClassificationOutput {
fn sync(self) -> Self {
self.loss.device().flush();
MultiLabelClassificationOutput {
output: self.output.no_grad(),
loss: self.loss.no_grad(),
targets: self.targets,
}
}
}
impl Adaptor<HammingScoreInput> for MultiLabelClassificationOutput {
fn adapt(&self) -> HammingScoreInput {
HammingScoreInput::new(self.output.clone(), self.targets.clone())
}
}
impl Adaptor<LossInput> for MultiLabelClassificationOutput {
fn adapt(&self) -> LossInput {
LossInput::new(self.loss.clone())
}
}
impl Adaptor<ConfusionStatsInput> for MultiLabelClassificationOutput {
fn adapt(&self) -> ConfusionStatsInput {
ConfusionStatsInput::new(self.output.clone(), self.targets.clone().bool())
}
}