This is an automated email from the ASF dual-hosted git repository.
CRZbulabula pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/iotdb.git
The following commit(s) were added to refs/heads/master by this push:
new b1a046211b1 [AINode] Replace DiT-derived Sundial conditioning head
(#18616)
b1a046211b1 is described below
commit b1a046211b1ade9f1bcd91cb6b323f4f44b3546b
Author: Yongzao <[email protected]>
AuthorDate: Thu Sep 10 21:28:45 2026 +0800
[AINode] Replace DiT-derived Sundial conditioning head (#18616)
---
.../iotdb/ainode/core/model/sundial/flow_loss.py | 78 ++++++-------
iotdb-core/ainode/tests/test_sundial_flow_loss.py | 123 +++++++++++++++++++++
2 files changed, 159 insertions(+), 42 deletions(-)
diff --git a/iotdb-core/ainode/iotdb/ainode/core/model/sundial/flow_loss.py
b/iotdb-core/ainode/iotdb/ainode/core/model/sundial/flow_loss.py
index b3fe95dbe2d..76a49431edf 100644
--- a/iotdb-core/ainode/iotdb/ainode/core/model/sundial/flow_loss.py
+++ b/iotdb-core/ainode/iotdb/ainode/core/model/sundial/flow_loss.py
@@ -70,10 +70,6 @@ class FlowLoss(nn.Module):
return x
-def modulate(x, shift, scale):
- return x * (1 + scale) + shift
-
-
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
@@ -119,11 +115,8 @@ class TimestepEmbedder(nn.Module):
return t_emb
-class ResBlock(nn.Module):
- """
- A residual block that can optionally change the number of channels.
- :param channels: the number of input channels.
- """
+class ConditionalResidualBlock(nn.Module):
+ """A residual MLP controlled by a per-sample conditioning vector."""
def __init__(self, channels):
super().__init__()
@@ -140,17 +133,21 @@ class ResBlock(nn.Module):
nn.SiLU(), nn.Linear(channels, 3 * channels, bias=True)
)
- def forward(self, x, y):
- shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(y).chunk(3,
dim=-1)
- h = modulate(self.in_ln(x), shift_mlp, scale_mlp)
- h = self.mlp(h)
- return x + gate_mlp * h
+ def reset_conditioning_parameters(self):
+ nn.init.zeros_(self.adaLN_modulation[-1].weight)
+ nn.init.zeros_(self.adaLN_modulation[-1].bias)
+ def forward(self, features, condition):
+ offset, gain_delta, update_scale = torch.tensor_split(
+ self.adaLN_modulation(condition), 3, dim=-1
+ )
+ conditioned = torch.addcmul(offset, self.in_ln(features),
gain_delta.add(1))
+ update = self.mlp(conditioned)
+ return torch.addcmul(features, update, update_scale)
-class FinalLayer(nn.Module):
- """
- The final layer adopted from DiT.
- """
+
+class ConditionalOutputProjection(nn.Module):
+ """Map condition-normalized features to the requested output width."""
def __init__(self, model_channels, out_channels):
super().__init__()
@@ -162,11 +159,20 @@ class FinalLayer(nn.Module):
nn.SiLU(), nn.Linear(model_channels, 2 * model_channels, bias=True)
)
- def forward(self, x, c):
- shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1)
- x = modulate(self.norm_final(x), shift, scale)
- x = self.linear(x)
- return x
+ def reset_conditioning_parameters(self):
+ nn.init.zeros_(self.adaLN_modulation[-1].weight)
+ nn.init.zeros_(self.adaLN_modulation[-1].bias)
+ nn.init.zeros_(self.linear.weight)
+ nn.init.zeros_(self.linear.bias)
+
+ def forward(self, features, condition):
+ offset, gain_delta = torch.tensor_split(
+ self.adaLN_modulation(condition), 2, dim=-1
+ )
+ conditioned = torch.addcmul(
+ offset, self.norm_final(features), gain_delta.add(1)
+ )
+ return self.linear(conditioned)
class SimpleMLPAdaLN(nn.Module):
@@ -199,16 +205,10 @@ class SimpleMLPAdaLN(nn.Module):
self.input_proj = nn.Linear(in_channels, model_channels)
- res_blocks = []
- for i in range(num_res_blocks):
- res_blocks.append(
- ResBlock(
- model_channels,
- )
- )
-
- self.res_blocks = nn.ModuleList(res_blocks)
- self.final_layer = FinalLayer(model_channels, out_channels)
+ self.res_blocks = nn.ModuleList(
+ ConditionalResidualBlock(model_channels) for _ in
range(num_res_blocks)
+ )
+ self.final_layer = ConditionalOutputProjection(model_channels,
out_channels)
self.initialize_weights()
@@ -225,16 +225,10 @@ class SimpleMLPAdaLN(nn.Module):
nn.init.normal_(self.time_embed.mlp[0].weight, std=0.02)
nn.init.normal_(self.time_embed.mlp[2].weight, std=0.02)
- # Zero-out adaLN modulation layers
for block in self.res_blocks:
- nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
- nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
-
- # Zero-out output layers
- nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
- nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
- nn.init.constant_(self.final_layer.linear.weight, 0)
- nn.init.constant_(self.final_layer.linear.bias, 0)
+ block.reset_conditioning_parameters()
+
+ self.final_layer.reset_conditioning_parameters()
def forward(self, x, t, c):
"""
diff --git a/iotdb-core/ainode/tests/test_sundial_flow_loss.py
b/iotdb-core/ainode/tests/test_sundial_flow_loss.py
new file mode 100644
index 00000000000..cb81f5f7a67
--- /dev/null
+++ b/iotdb-core/ainode/tests/test_sundial_flow_loss.py
@@ -0,0 +1,123 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+#
+
+import unittest
+
+import torch
+
+from iotdb.ainode.core.model.sundial.flow_loss import (
+ ConditionalOutputProjection,
+ ConditionalResidualBlock,
+ FlowLoss,
+ SimpleMLPAdaLN,
+)
+
+
+class SundialFlowLossTest(unittest.TestCase):
+ def create_network(self):
+ return SimpleMLPAdaLN(
+ in_channels=4,
+ model_channels=8,
+ out_channels=4,
+ z_channels=6,
+ num_res_blocks=2,
+ )
+
+ def test_checkpoint_parameter_contract(self):
+ network = self.create_network()
+ expected_shapes = {
+ "time_embed.mlp.0.weight": (8, 256),
+ "time_embed.mlp.0.bias": (8,),
+ "time_embed.mlp.2.weight": (8, 8),
+ "time_embed.mlp.2.bias": (8,),
+ "cond_embed.weight": (8, 6),
+ "cond_embed.bias": (8,),
+ "input_proj.weight": (8, 4),
+ "input_proj.bias": (8,),
+ "res_blocks.0.in_ln.weight": (8,),
+ "res_blocks.0.in_ln.bias": (8,),
+ "res_blocks.0.mlp.0.weight": (8, 8),
+ "res_blocks.0.mlp.0.bias": (8,),
+ "res_blocks.0.mlp.2.weight": (8, 8),
+ "res_blocks.0.mlp.2.bias": (8,),
+ "res_blocks.0.adaLN_modulation.1.weight": (24, 8),
+ "res_blocks.0.adaLN_modulation.1.bias": (24,),
+ "res_blocks.1.in_ln.weight": (8,),
+ "res_blocks.1.in_ln.bias": (8,),
+ "res_blocks.1.mlp.0.weight": (8, 8),
+ "res_blocks.1.mlp.0.bias": (8,),
+ "res_blocks.1.mlp.2.weight": (8, 8),
+ "res_blocks.1.mlp.2.bias": (8,),
+ "res_blocks.1.adaLN_modulation.1.weight": (24, 8),
+ "res_blocks.1.adaLN_modulation.1.bias": (24,),
+ "final_layer.linear.weight": (4, 8),
+ "final_layer.linear.bias": (4,),
+ "final_layer.adaLN_modulation.1.weight": (16, 8),
+ "final_layer.adaLN_modulation.1.bias": (16,),
+ }
+
+ actual_shapes = {
+ key: tuple(value.shape) for key, value in
network.state_dict().items()
+ }
+ self.assertEqual(expected_shapes, actual_shapes)
+
+ restored = self.create_network()
+ result = restored.load_state_dict(network.state_dict(), strict=True)
+ self.assertEqual([], result.missing_keys)
+ self.assertEqual([], result.unexpected_keys)
+
+ def test_conditioned_layers_propagate_gradients(self):
+ features = torch.randn(3, 8, requires_grad=True)
+ condition = torch.randn(3, 8, requires_grad=True)
+ block = ConditionalResidualBlock(8)
+ projection = ConditionalOutputProjection(8, 4)
+
+ output = projection(block(features, condition), condition)
+
+ self.assertEqual((3, 4), tuple(output.shape))
+ output.square().mean().backward()
+ self.assertIsNotNone(features.grad)
+ self.assertIsNotNone(condition.grad)
+
+ def test_network_starts_with_zero_output(self):
+ network = self.create_network()
+
+ output = network(
+ torch.randn(3, 4),
+ torch.tensor([0.0, 500.0, 999.0]),
+ torch.randn(3, 6),
+ )
+
+ torch.testing.assert_close(output, torch.zeros_like(output))
+
+ def test_sample_shape(self):
+ flow_loss = FlowLoss(
+ target_channels=4,
+ z_channels=6,
+ depth=2,
+ width=8,
+ num_sampling_steps=2,
+ )
+
+ samples = flow_loss.sample(torch.randn(2, 6), num_samples=3)
+
+ self.assertEqual((2, 3, 4), tuple(samples.shape))
+
+
+if __name__ == "__main__":
+ unittest.main()