Skip to content

Commit 57fe4aa

Browse files
committed
Conditioning via AdaLN
1 parent 0308cf7 commit 57fe4aa

7 files changed

Lines changed: 286 additions & 53 deletions

File tree

scripts/data.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -285,6 +285,8 @@ def compute_symmetry_sites(
285285
symmetry_dataset["composition"] = symmetry_dataset.apply(get_composition_from_symmetry_sites, axis=1)
286286
if "formation_energy_per_atom" in dataset.columns:
287287
symmetry_dataset['formation_energy_per_atom'] = dataset['formation_energy_per_atom']
288+
if "energy_above_hull" in dataset.columns:
289+
symmetry_dataset['energy_above_hull'] = dataset['energy_above_hull']
288290
if "band_gap" in dataset.columns:
289291
symmetry_dataset['band_gap'] = dataset['band_gap']
290292
if "log_klat" in dataset.columns:

src/wyckoff_transformer/cascade/dataset.py

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -170,7 +170,8 @@ def __init__(
170170
start_dtype: torch.dtype = torch.int64,
171171
device: str = "cpu",
172172
augmented_storage_device: Optional[str] = None,
173-
target_name = None):
173+
target_name = None,
174+
extra_fields: Optional[List[str]] = None):
174175
"""
175176
Args:
176177
data (dict[str, Tensor | List[Tensor]]):
@@ -234,6 +235,9 @@ def __init__(
234235
target_name (Optional[str], optional):
235236
The key in `data` for the target variable, if any. The corresponding
236237
value `data[target_name]` should be a `torch.Tensor`. Defaults to None.
238+
extra_fields (Optional[List[str]], optional):
239+
Optional list of additional keys in `data` to copy into `self.data`
240+
(e.g. condition features used for AdaLN). Defaults to None.
237241
"""
238242
if augmented_fields is None:
239243
augmented_fields = []
@@ -250,7 +254,11 @@ def __init__(
250254
self.cascade_index_from_field = {name: i for i, name in enumerate(cascade_order)}
251255
self.augmented_fields = augmented_fields
252256
self.data = {name: data[name].type(dtype).to(device) for name in cascade_order if name not in augmented_fields}
253-
self.max_sequence_length = next(iter(self.data.values())).size(1)
257+
if extra_fields:
258+
for k in extra_fields:
259+
v = data[k]
260+
self.data[k] = v.to(device) if hasattr(v, "to") else v
261+
self.max_sequence_length = data[cascade_order[0]].size(1)
254262
self.masks = {name: torch.tensor(masks[name], dtype=dtype, device=device) for name in cascade_order}
255263
self.pads = {name: torch.tensor(pads[name], dtype=dtype, device=device) for name in cascade_order}
256264
self.stops = {name: torch.tensor(stops[name], dtype=dtype, device=device) for name in cascade_order}
@@ -346,7 +354,8 @@ def get_augmentation(
346354
def get_masked_cascade_data(
347355
self,
348356
known_seq_len: int,
349-
known_cascade_len: int):
357+
known_cascade_len: int,
358+
return_chosen_indices: bool = False):
350359

351360
# assert known_seq_len < data.shape[1]
352361
# assert known_seq_len >= 0
@@ -372,6 +381,9 @@ def get_masked_cascade_data(
372381
res.append(torch.cat([
373382
cascade_vector[:, :known_seq_len],
374383
self.masks[name].expand(cascade_vector.size(0), 1)], dim=1))
384+
385+
if return_chosen_indices:
386+
return self.start_tokens[target_is_viable], res, target[target_is_viable], target_is_viable
375387
return self.start_tokens[target_is_viable], res, target[target_is_viable]
376388

377389

src/wyckoff_transformer/cascade/model.py

Lines changed: 68 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,53 @@
1111

1212
logger = logging.getLogger(__name__)
1313

14+
15+
class AdaLNTransformerEncoderLayer(TransformerEncoderLayer):
16+
def __init__(self, d_model, nhead, condition_dim=1, layer_norm_eps=1e-5, **kwargs):
17+
super().__init__(d_model, nhead, layer_norm_eps=layer_norm_eps, **kwargs)
18+
# AdaLN provides the affine; strip it from the parent's LayerNorms.
19+
self.norm1 = nn.LayerNorm(d_model, eps=layer_norm_eps, elementwise_affine=False)
20+
self.norm2 = nn.LayerNorm(d_model, eps=layer_norm_eps, elementwise_affine=False)
21+
# DiT-style (1 + gamma) modulation; zero-init weights/biases => identity at init.
22+
self.adaLN_modulation1 = nn.Linear(condition_dim, 2 * d_model)
23+
self.adaLN_modulation2 = nn.Linear(condition_dim, 2 * d_model)
24+
for m in (self.adaLN_modulation1, self.adaLN_modulation2):
25+
nn.init.zeros_(m.weight)
26+
nn.init.zeros_(m.bias)
27+
28+
def forward(self, src, cond, src_mask=None, src_key_padding_mask=None, is_causal=False):
29+
gamma1, beta1 = self.adaLN_modulation1(cond).chunk(2, dim=-1)
30+
gamma2, beta2 = self.adaLN_modulation2(cond).chunk(2, dim=-1)
31+
32+
x = src
33+
if self.norm_first:
34+
x_norm = self.norm1(x) * (1 + gamma1.unsqueeze(1)) + beta1.unsqueeze(1)
35+
x = x + self.dropout1(self.self_attn(x_norm, x_norm, x_norm, attn_mask=src_mask,
36+
key_padding_mask=src_key_padding_mask, need_weights=False, is_causal=is_causal)[0])
37+
x_norm2 = self.norm2(x) * (1 + gamma2.unsqueeze(1)) + beta2.unsqueeze(1)
38+
x = x + self.dropout2(self.linear2(self.dropout(self.activation(self.linear1(x_norm2)))))
39+
else:
40+
x2 = self.self_attn(x, x, x, attn_mask=src_mask,
41+
key_padding_mask=src_key_padding_mask, need_weights=False, is_causal=is_causal)[0]
42+
x = x + self.dropout1(x2)
43+
x = self.norm1(x) * (1 + gamma1.unsqueeze(1)) + beta1.unsqueeze(1)
44+
45+
x2 = self.linear2(self.dropout(self.activation(self.linear1(x))))
46+
x = x + self.dropout2(x2)
47+
x = self.norm2(x) * (1 + gamma2.unsqueeze(1)) + beta2.unsqueeze(1)
48+
return x
49+
50+
class AdaLNTransformerEncoder(TransformerEncoder):
51+
"""`nn.TransformerEncoder` that threads a conditioning tensor to each AdaLN layer."""
52+
def forward(self, src, cond, mask=None, src_key_padding_mask=None, is_causal=False):
53+
output = src
54+
for mod in self.layers:
55+
output = mod(output, cond=cond, src_mask=mask,
56+
src_key_padding_mask=src_key_padding_mask, is_causal=is_causal)
57+
if self.norm is not None:
58+
output = self.norm(output)
59+
return output
60+
1461
class SpecialEmbedding(torch.nn.Module):
1562
ScalarPassThrough = 0
1663
VectorPassThrough = 1
@@ -199,7 +246,8 @@ def __init__(self,
199246
aggregation_weight: Optional[int] = None,
200247
emebdding_dropout: Optional[float] = None,
201248
prediction_perceptron_dropout: Optional[float] = None,
202-
concat_start_to_prediction_input_embedding_dim: Optional[int] = None):
249+
concat_start_to_prediction_input_embedding_dim: Optional[int] = None,
250+
condition_dim: Optional[int] = None):
203251
"""
204252
Expects tokens in the following format:
205253
START_k -> [] -> STOP -> PAD
@@ -250,8 +298,16 @@ def __init__(self,
250298
if "nhead" in TransformerEncoderLayer_args and self.d_model % TransformerEncoderLayer_args["nhead"]:
251299
logger.warning("d_model is not divisible by nhead, padding to the next multiple")
252300
self.d_model += TransformerEncoderLayer_args["nhead"] - self.d_model % TransformerEncoderLayer_args["nhead"]
253-
self.encoder_layers = TransformerEncoderLayer(self.d_model, batch_first=True, **TransformerEncoderLayer_args)
254-
self.transformer_encoder = TransformerEncoder(self.encoder_layers, **TransformerEncoder_args)
301+
302+
self.condition_dim = condition_dim
303+
if condition_dim is not None:
304+
self.encoder_layers = AdaLNTransformerEncoderLayer(
305+
self.d_model, batch_first=True, condition_dim=condition_dim, **TransformerEncoderLayer_args)
306+
self.transformer_encoder = AdaLNTransformerEncoder(self.encoder_layers, **TransformerEncoder_args)
307+
else:
308+
self.encoder_layers = TransformerEncoderLayer(self.d_model, batch_first=True, **TransformerEncoderLayer_args)
309+
self.transformer_encoder = TransformerEncoder(self.encoder_layers, **TransformerEncoder_args)
310+
255311
self.start_type = start_type
256312
if start_type == "categorial":
257313
self.start_embedding = nn.Embedding(n_start, self.d_model)
@@ -347,7 +403,8 @@ def forward(self,
347403
start: Tensor,
348404
cascade: List[Tensor],
349405
padding_mask: Tensor|None,
350-
prediction_head: int|None) -> Tensor:
406+
prediction_head: int|None,
407+
cond: Tensor|None = None) -> Tensor:
351408
"""
352409
Arguments:
353410
start: Tensor of shape ``[batch_size]`` with the start token.
@@ -356,6 +413,7 @@ def forward(self,
356413
prediction_head: Index of the prediction head to use. If None, use the only one. The
357414
model works in two stages. Firstly, a vector is prepared wih Encoder and
358415
various tweaks. Then, the vector is passed to a perceptron aka prediction head.
416+
cond: Tensor of shape ``[batch_size, condition_dim]`` with the conditioning vector for AdaLN.
359417
Returns:
360418
Tensor of shape ``[batch_size, seq_len, output_dim]`` with the predictions.
361419
"""
@@ -377,7 +435,12 @@ def forward(self,
377435
data = torch.cat([self.start_embedding(start).unsqueeze(1), cascade_embedding], dim=1)
378436
logger.debug("Data size: %s", data.size())
379437
logger.debug("Padding mask size: %s", padding_mask.size() if padding_mask is not None else "None")
380-
transformer_output = self.transformer_encoder(data, src_key_padding_mask=padding_mask)
438+
if getattr(self, "condition_dim", None) is not None:
439+
if cond is None:
440+
raise ValueError("condition_dim is set but cond is not provided")
441+
transformer_output = self.transformer_encoder(data, src_key_padding_mask=padding_mask, cond=cond)
442+
else:
443+
transformer_output = self.transformer_encoder(data, src_key_padding_mask=padding_mask)
381444

382445
logging.debug("Transformer output size: %s", transformer_output.size())
383446
if self.aggregate_after_encoder:

src/wyckoff_transformer/generator.py

Lines changed: 14 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -63,13 +63,16 @@ def __init__(self,
6363
self.token_engineers = token_engineers
6464

6565

66-
def calibrate(self, dataset: AugmentedCascadeDataset, calibration_element_count_threshold: int = 100):
66+
def calibrate(self, dataset: AugmentedCascadeDataset, calibration_element_count_threshold: int = 100, condition_feature: Optional[str] = None):
6767
"""
6868
The calibraiton is going to be per cascade field.
6969
We will generate p_predicted and p_true for each cascade field for each
7070
known sequence length.
7171
"""
7272
assert dataset.cascade_order == self.cascade_order
73+
74+
full_cond = dataset.data[condition_feature] if condition_feature is not None else None
75+
7376
with torch.no_grad():
7477
self.model.eval()
7578
self.calibrators = []
@@ -84,14 +87,17 @@ def calibrate(self, dataset: AugmentedCascadeDataset, calibration_element_count_
8487
logging.info("Calibrating cascade field %s", cascade_name)
8588
for known_seq_len in range(dataset.max_sequence_length):
8689
if known_cascade_len == 0:
87-
start_tokens, masked_data, target = dataset.get_masked_multiclass_cascade_data(
90+
start_tokens, masked_data, target, chosen_indices = dataset.get_masked_multiclass_cascade_data(
8891
known_seq_len, known_cascade_len,
89-
target_type=TargetClass.NextToken, multiclass_target=True)
92+
target_type=TargetClass.NextToken, multiclass_target=True,
93+
return_chosen_indices=True)
9094
else:
91-
start_tokens, masked_data, target = dataset.get_masked_multiclass_cascade_data(
95+
start_tokens, masked_data, target, chosen_indices = dataset.get_masked_multiclass_cascade_data(
9296
known_seq_len, known_cascade_len,
93-
target_type=TargetClass.NextToken, multiclass_target=False)
94-
model_output = self.model(start_tokens, masked_data, None, known_cascade_len)
97+
target_type=TargetClass.NextToken, multiclass_target=False,
98+
return_chosen_indices=True)
99+
iter_cond = full_cond[chosen_indices] if full_cond is not None else None
100+
model_output = self.model(start_tokens, masked_data, None, known_cascade_len, cond=iter_cond)
95101
# Enought data for separate calibration
96102
if target.size(0) >= calibration_element_count_threshold:
97103
# If model_output and target are on different devices, not our problem
@@ -125,6 +131,7 @@ def generate_tensors(
125131
max_length: Optional[int] = None,
126132
elements_vocab: Optional[Dict] = None,
127133
delimiter: str = "-",
134+
cond: Optional[Tensor] = None
128135
) -> List[Tensor] | Tuple[List[Tensor], List[float], List[float]]:
129136
"""
130137
Generates a sequence of tokens.
@@ -245,7 +252,7 @@ def _parse_string_to_ids(s: str) -> set:
245252
if self.cascade_is_target.get(cascade_name, False):
246253
# +1 for MASK
247254
this_generation_input = [generated_cascade[:, :known_seq_len + 1] for generated_cascade in generated]
248-
logits = self.model(start, this_generation_input, None, known_cascade_len)
255+
logits = self.model(start, this_generation_input, None, known_cascade_len, cond=cond)
249256
if self.calibrators is not None:
250257
if known_seq_len < len(self.calibrators[known_cascade_len]):
251258
logits = self.calibrators[known_cascade_len][known_seq_len](logits)

0 commit comments

Comments
 (0)