1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
use crate::{Tensor, kind::Autodiff};
#[cfg(feature = "autodiff")]
use crate::ops::{BridgeKind, BridgeTensor};
#[cfg(feature = "autodiff")]
use burn_backend::AutodiffBackend;
#[cfg(feature = "autodiff")]
use burn_dispatch::Dispatch;
#[cfg(feature = "autodiff")]
use burn_dispatch::GradientCheckpointingStrategy;
#[cfg(feature = "autodiff")]
type AutodiffGradients = <Dispatch as AutodiffBackend>::Gradients;
// Aligned, type-erased storage for `AutodiffGradients`. See `crate::macros`
// for why this indirection exists.
#[cfg(feature = "autodiff")]
burn_std::obfuscate!(
type: AutodiffGradients,
module: gradients_opaque,
derives: [Send]
);
/// Gradients container used during the backward pass.
#[cfg(feature = "autodiff")]
pub struct Gradients {
blob: gradients_opaque::Opaque,
}
#[cfg(feature = "autodiff")]
impl Gradients {
/// Crate-internal constructor wrapping the dispatch-level gradients.
pub(crate) fn from_inner(inner: AutodiffGradients) -> Self {
Self {
blob: gradients_opaque::Opaque::new(inner),
}
}
/// Crate-internal borrow of the underlying gradients container.
pub(crate) fn as_inner(&self) -> &AutodiffGradients {
self.blob.as_ref()
}
/// Crate-internal mutable borrow of the underlying gradients container.
pub(crate) fn as_inner_mut(&mut self) -> &mut AutodiffGradients {
self.blob.as_mut()
}
}
#[cfg(feature = "autodiff")]
impl<const D: usize> Tensor<D> {
/// Backward pass of the tensor.
pub fn backward(&self) -> Gradients {
backward_impl(&self.primitive)
}
/// Get the gradients of a tensor if it exist.
///
/// Returns a new reference to the same tensor. Therefore the same grad tensor can
/// be accessed multiple times. If you only need to get the gradients one time,
/// consider using [grad_remove](Tensor::grad_remove) for better performance.
pub fn grad(&self, grads: &Gradients) -> Option<Tensor<D>> {
grad_impl(&self.primitive, grads).map(Tensor::new)
}
/// Remove the grad tensor from the [grads](AutodiffBackend::Gradients) struct returning the result.
pub fn grad_remove(&self, grads: &mut Gradients) -> Option<Tensor<D>> {
grad_remove_impl(&self.primitive, grads).map(Tensor::new)
}
/// Replace the grad tensor from the [grads](AutodiffBackend::Gradients) struct with the provided
/// gradient.
pub fn grad_replace(&self, grads: &mut Gradients, grad: Tensor<D>) {
grad_replace_impl(&self.primitive, grads, grad.primitive)
}
}
#[cfg(feature = "autodiff")]
fn backward_impl(p: &BridgeTensor) -> Gradients {
Gradients::from_inner(Dispatch::backward(p.clone().into_float()))
}
#[cfg(feature = "autodiff")]
fn grad_impl(p: &BridgeTensor, grads: &Gradients) -> Option<BridgeTensor> {
// A non-float tensor — a packed base included — records no tape, so there
// is no gradient to look up.
Dispatch::grad(p.try_as_float()?, grads.as_inner()).map(BridgeTensor::float)
}
#[cfg(feature = "autodiff")]
fn grad_remove_impl(p: &BridgeTensor, grads: &mut Gradients) -> Option<BridgeTensor> {
Dispatch::grad_remove(p.try_as_float()?, grads.as_inner_mut()).map(BridgeTensor::float)
}
#[cfg(feature = "autodiff")]
fn grad_replace_impl(p: &BridgeTensor, grads: &mut Gradients, grad: BridgeTensor) {
Dispatch::grad_replace(p.as_float(), grads.as_inner_mut(), grad.into_float())
}
impl<const D: usize, K: Autodiff> Tensor<D, K> {
/// Returns the inner tensor without the autodiff information.
pub fn inner(self) -> Tensor<D, K> {
Tensor::new(K::inner(self.primitive))
}
/// Take the tensor off the autodiff backend, dropping any graph reference
/// it carries. A tensor that is not on an autodiff backend is returned as
/// is.
///
/// Unlike [detach](Tensor::detach), which severs the tensor from the graph
/// but leaves it on the autodiff backend, the result lives on the inner
/// backend: later operations pay no autodiff dispatch and can never be
/// recorded. And unlike [inner](Self::inner), which panics on a tensor
/// with no autodiff wrapper, this is safe to call anywhere — batchers,
/// metric pipelines, anything that must guarantee a tensor is off the
/// tape without knowing where it came from.
pub fn no_grad(self) -> Self {
if self.device().is_autodiff() {
self.inner()
} else {
self
}
}
/// Convert a tensor to the autodiff backend.
///
/// # Arguments
///
/// * `inner` - The tensor to convert.
///
/// # Returns
///
/// The tensor converted to the autodiff backend.
pub fn from_inner(inner: Tensor<D, K>) -> Self {
Self::new(K::from_inner(inner.primitive))
}
/// Sets the autodiff checkpointing strategy carried by this tensor.
///
/// The strategy is normally derived from the device the tensor was created on (see
/// [`Device::gradient_checkpointing`](crate::Device::gradient_checkpointing)); this
/// method overrides it for a single tensor. A tensor carrying a strategy is treated
/// as tracked by autodiff, so this also marks an inner-backend tensor for tracking.
///
/// # Panics
///
/// Operations combining tensors that carry different strategies panic; make sure all
/// operands share the same one.
#[cfg(feature = "autodiff")]
pub fn with_gradient_checkpointing_strategy(
self,
strategy: GradientCheckpointingStrategy,
) -> Self {
let (kind, mut tensor) = self.primitive.into_parts();
tensor.checkpointing = Some(strategy);
Self::new(match kind {
BridgeKind::Bool => BridgeTensor::bool(tensor),
BridgeKind::Int => BridgeTensor::int(tensor),
BridgeKind::Float => BridgeTensor::float(tensor),
BridgeKind::QFloat => BridgeTensor::qfloat(tensor),
})
}
}
// TODO: a lot of the `tensor.inner` / `Tensor::from_inner(...)` are actually scoped to perform some operations
// so it might be cleaner and easier to manage the device etc. if we provide a method to scope the autodiff?