Views
No views yet
1from transformers import (
2 ViTForImageClassification,
3 pipeline,
4 AutoImageProcessor,
5 ViTConfig,
6 ViTModel,
7)
8
9from transformers.modeling_outputs import (
10 ImageClassifierOutput,
11 BaseModelOutputWithPooling,
12)
13
14from PIL import Image
15import torch
16from torch import nn
17from typing import Optional, Union, Tuple
18
19
20class CustomViTModel(ViTModel):
21 def forward(
22 self,
23 pixel_values: Optional[torch.Tensor] = None,
24 bool_masked_pos: Optional[torch.BoolTensor] = None,
25 head_mask: Optional[torch.Tensor] = None,
26 output_attentions: Optional[bool] = None,
27 output_hidden_states: Optional[bool] = None,
28 interpolate_pos_encoding: Optional[bool] = None,
29 return_dict: Optional[bool] = None,
30 ) -> Union[Tuple, BaseModelOutputWithPooling]:
31 r"""
32 bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, num_patches)`, *optional*):
33 Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).
34 """
35 output_attentions = (
36 output_attentions
37 if output_attentions is not None
38 else self.config.output_attentions
39 )
40 output_hidden_states = (
41 output_hidden_states
42 if output_hidden_states is not None
43 else self.config.output_hidden_states
44 )
45 return_dict = (
46 return_dict if return_dict is not None else self.config.use_return_dict
47 )
48
49 if pixel_values is None:
50 raise ValueError("You have to specify pixel_values")
51
52 # Prepare head mask if needed
53 # 1.0 in head_mask indicate we keep the head
54 # attention_probs has shape bsz x n_heads x N x N
55 # input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]
56 # and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
57 head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)
58
59 # TODO: maybe have a cleaner way to cast the input (from `ImageProcessor` side?)
60 expected_dtype = self.embeddings.patch_embeddings.projection.weight.dtype
61 if pixel_values.dtype != expected_dtype:
62 pixel_values = pixel_values.to(expected_dtype)
63
64 embedding_output = self.embeddings(
65 pixel_values,
66 bool_masked_pos=bool_masked_pos,
67 interpolate_pos_encoding=interpolate_pos_encoding,
68 )
69
70 encoder_outputs = self.encoder(
71 embedding_output,
72 head_mask=head_mask,
73 output_attentions=output_attentions,
74 output_hidden_states=output_hidden_states,
75 return_dict=return_dict,
76 )
77 sequence_output = encoder_outputs[0]
78 sequence_output = sequence_output[:, 1:, :].mean(dim=1)
79
80 sequence_output = self.layernorm(sequence_output)
81 pooled_output = (
82 self.pooler(sequence_output) if self.pooler is not None else None
83 )
84
85 if not return_dict:
86 head_outputs = (
87 (sequence_output, pooled_output)
88 if pooled_output is not None
89 else (sequence_output,)
90 )
91 return head_outputs + encoder_outputs[1:]
92
93 return BaseModelOutputWithPooling(
94 last_hidden_state=sequence_output,
95 pooler_output=pooled_output,
96 hidden_states=encoder_outputs.hidden_states,
97 attentions=encoder_outputs.attentions,
98 )
99
100
101class CustomViTForImageClassification(ViTForImageClassification):
102 def __init__(self, config: ViTConfig) -> None:
103 super().__init__(config)
104
105 self.num_labels = config.num_labels
106 self.vit = CustomViTModel(config, add_pooling_layer=False)
107
108 # Classifier head
109 self.classifier = (
110 nn.Linear(config.hidden_size, config.num_labels)
111 if config.num_labels > 0
112 else nn.Identity()
113 )
114
115 # Initialize weights and apply final processing
116 self.post_init()
117
118 def forward(
119 self,
120 pixel_values: Optional[torch.Tensor] = None,
121 head_mask: Optional[torch.Tensor] = None,
122 labels: Optional[torch.Tensor] = None,
123 output_attentions: Optional[bool] = None,
124 output_hidden_states: Optional[bool] = None,
125 interpolate_pos_encoding: Optional[bool] = None,
126 return_dict: Optional[bool] = None,
127 ) -> Union[tuple, ImageClassifierOutput]:
128 r"""
129 labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
130 Labels for computing the image classification/regression loss. Indices should be in `[0, ...,
131 config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
132 `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
133 """
134 return_dict = (
135 return_dict if return_dict is not None else self.config.use_return_dict
136 )
137
138 outputs = self.vit(
139 pixel_values,
140 head_mask=head_mask,
141 output_attentions=output_attentions,
142 output_hidden_states=output_hidden_states,
143 interpolate_pos_encoding=interpolate_pos_encoding,
144 return_dict=return_dict,
145 )
146
147 sequence_output = outputs[0]
148
149 logits = self.classifier(sequence_output)
150
151 loss = None
152
153 return ImageClassifierOutput(
154 loss=loss,
155 logits=logits,
156 hidden_states=outputs.hidden_states,
157 attentions=outputs.attentions,
158 )
159
160if __name__ == "__main__":
161
162 model = CustomViTForImageClassification.from_pretrained("vesteinn/vit-mae-inat21")
163 image_processor = AutoImageProcessor.from_pretrained("vesteinn/vit-mae-inat21")
164
165 classifier = pipeline(
166 "image-classification", model=model, image_processor=image_processor
167 )