Coverage for transformer_lens/model_bridge/generalized_components/bloom_block.py: 36%

42 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""BLOOM-specific block bridge component. 

2 

3BLOOM blocks require special arguments (alibi, attention_mask, residual) that standard 

4BlockBridge doesn't handle. This custom component generates and passes these arguments. 

5""" 

6from typing import Any, Dict, Optional 

7 

8import torch 

9 

10from transformer_lens.model_bridge.generalized_components.alibi_utils import ( 

11 build_alibi_tensor as _build_alibi_tensor, 

12) 

13from transformer_lens.model_bridge.generalized_components.base import ( 

14 GeneralizedComponent, 

15) 

16from transformer_lens.model_bridge.generalized_components.block import BlockBridge 

17 

18 

19class BloomBlockBridge(BlockBridge): 

20 """Block bridge for BLOOM models that handles ALiBi positional encoding. 

21 

22 BLOOM uses ALiBi (Attention with Linear Biases) instead of standard positional 

23 embeddings. This requires generating an alibi tensor and passing it to each block 

24 along with the attention_mask. 

25 """ 

26 

27 def __init__( 

28 self, 

29 name: str, 

30 config: Optional[Any] = None, 

31 submodules: Optional[Dict[str, GeneralizedComponent]] = None, 

32 hook_alias_overrides: Optional[Dict[str, str]] = None, 

33 ): 

34 """Initialize the BLOOM block bridge. 

35 

36 Args: 

37 name: The name of the component in the model 

38 config: Model configuration (used to get n_heads for ALiBi) 

39 submodules: Dictionary of submodules to register 

40 hook_alias_overrides: Optional dictionary to override default hook aliases 

41 """ 

42 super().__init__(name, config, submodules, hook_alias_overrides) 

43 self.config = config 

44 

45 @staticmethod 

46 def build_alibi_tensor( 

47 attention_mask: torch.Tensor, num_heads: int, dtype: torch.dtype 

48 ) -> torch.Tensor: 

49 """Build ALiBi tensor for attention biasing. 

50 

51 Delegates to the shared ALiBi utility in alibi_utils.py. 

52 

53 Args: 

54 attention_mask: Attention mask of shape [batch_size, seq_length] 

55 num_heads: Number of attention heads 

56 dtype: Data type for the tensor 

57 

58 Returns: 

59 ALiBi tensor of shape [batch_size, num_heads, 1, seq_length]. 

60 """ 

61 return _build_alibi_tensor(attention_mask, num_heads, dtype) 

62 

63 def forward(self, *args: Any, **kwargs: Any) -> Any: 

64 """Forward pass through the BLOOM block. 

65 

66 BLOOM blocks require `alibi` and `attention_mask` arguments. If the HF model's 

67 BloomModel.forward() is being called, it will generate these and pass them through. 

68 If they're missing (e.g., when called standalone), we generate them here. 

69 

70 Args: 

71 *args: Positional arguments (first should be hidden_states) 

72 **kwargs: Keyword arguments 

73 

74 Returns: 

75 Output from the original BLOOM block 

76 """ 

77 if self.original_component is None: 77 ↛ 78line 77 didn't jump to line 78 because the condition on line 77 was never true

78 raise RuntimeError( 

79 f"Original component not set for {self.name}. Call set_original_component() first." 

80 ) 

81 

82 self._check_stop_at_layer(*args, **kwargs) 

83 args, kwargs = self._hook_input_hidden_states(args, kwargs) 

84 

85 # BLOOM blocks require 'alibi' and 'attention_mask' arguments. 

86 # If HF's BloomModel.forward() is calling us, these will already be present. 

87 # Only generate them if they're missing (e.g., standalone block testing). 

88 if "alibi" not in kwargs or kwargs["alibi"] is None: 88 ↛ 90line 88 didn't jump to line 90 because the condition on line 88 was never true

89 # Get hidden_states to determine shape 

90 if len(args) > 0 and isinstance(args[0], torch.Tensor): 

91 hidden_states = args[0] 

92 elif "hidden_states" in kwargs: 

93 hidden_states = kwargs["hidden_states"] 

94 else: 

95 raise ValueError("Could not find hidden_states in args or kwargs") 

96 

97 batch_size, seq_length, _ = hidden_states.shape 

98 device = hidden_states.device 

99 dtype = hidden_states.dtype 

100 

101 # Generate attention_mask if missing 

102 if "attention_mask" not in kwargs or kwargs["attention_mask"] is None: 

103 # Create default attention mask (all ones) 

104 attention_mask = torch.ones(batch_size, seq_length, dtype=torch.long, device=device) 

105 else: 

106 attention_mask = kwargs["attention_mask"] 

107 # Ensure it's 2D [batch, seq_length] for ALiBi generation 

108 if attention_mask.dim() == 4: 

109 # If 4D, we need 2D version for ALiBi generation 

110 # Extract the last row which tells us which positions are valid 

111 attention_mask_2d = attention_mask[:, 0, -1, :].long() 

112 elif attention_mask.dim() == 2: 

113 attention_mask_2d = attention_mask 

114 else: 

115 raise ValueError( 

116 f"Unexpected attention_mask dimensions: {attention_mask.dim()}" 

117 ) 

118 

119 # Generate ALiBi bias 

120 if self.config and hasattr(self.config, "n_heads"): 

121 num_heads = self.config.n_heads 

122 else: 

123 # Fallback: try to infer from model 

124 num_heads = 16 # BLOOM-560M has 16 heads 

125 

126 # Generate alibi — shared utility returns [batch, heads, 1, seq], 

127 # reshape to [batch*heads, 1, seq] to match HF's format for baddbmm. 

128 alibi = self.build_alibi_tensor(attention_mask_2d, num_heads, dtype) 

129 alibi = alibi.reshape(batch_size * num_heads, 1, seq_length) 

130 

131 # Add alibi to kwargs 

132 kwargs["alibi"] = alibi 

133 # else: alibi is already present from HF, don't overwrite it! 

134 

135 output = self.original_component(*args, **kwargs) 

136 return self._apply_output_hook(output, wrap_single_element=False)