MCPcopy Create free account
hub / github.com/cpystan/SD-VLM / initialize_vision_modules

Method initialize_vision_modules

llava/model/llava_arch.py:53–101  ·  view source on GitHub ↗
(self, model_args, fsdp=None)

Source from the content-addressed store, hash-verified

51 return vision_tower
52
53 def initialize_vision_modules(self, model_args, fsdp=None):
54 vision_tower = model_args.vision_tower
55 mm_vision_select_layer = model_args.mm_vision_select_layer
56 mm_vision_select_feature = model_args.mm_vision_select_feature
57 pretrain_mm_mlp_adapter = model_args.pretrain_mm_mlp_adapter
58 mm_patch_merge_type = model_args.mm_patch_merge_type
59
60 self.config.mm_vision_tower = vision_tower
61
62 if self.get_vision_tower() is None:
63 vision_tower = build_vision_tower(model_args)
64
65 if fsdp is not None and len(fsdp) > 0:
66 self.vision_tower = [vision_tower]
67 else:
68 self.vision_tower = vision_tower
69 else:
70 if fsdp is not None and len(fsdp) > 0:
71 vision_tower = self.vision_tower[0]
72 else:
73 vision_tower = self.vision_tower
74 vision_tower.load_model()
75
76 self.config.use_mm_proj = True
77 self.config.mm_projector_type = getattr(model_args, 'mm_projector_type', 'linear')
78 self.config.mm_hidden_size = vision_tower.hidden_size
79 self.config.mm_vision_select_layer = mm_vision_select_layer
80 self.config.mm_vision_select_feature = mm_vision_select_feature
81 self.config.mm_patch_merge_type = mm_patch_merge_type
82
83 if getattr(self, 'mm_projector', None) is None:
84 self.mm_projector = build_vision_projector(self.config)
85
86 if 'unpad' in mm_patch_merge_type:
87 embed_std = 1 / torch.sqrt(torch.tensor(self.config.hidden_size, dtype=self.dtype))
88 self.image_newline = nn.Parameter(
89 torch.randn(self.config.hidden_size, dtype=self.dtype) * embed_std
90 )
91 else:
92 # In case it is frozen by LoRA
93 for p in self.mm_projector.parameters():
94 p.requires_grad = True
95
96 if pretrain_mm_mlp_adapter is not None:
97 mm_projector_weights = torch.load(pretrain_mm_mlp_adapter, map_location='cpu')
98 def get_w(weights, keyword):
99 return {k.split(keyword + '.')[1]: v for k, v in weights.items() if keyword in k}
100
101 self.mm_projector.load_state_dict(get_w(mm_projector_weights, 'mm_projector'))
102
103
104def unpad_image(tensor, original_size):

Callers 1

trainFunction · 0.80

Calls 4

get_vision_towerMethod · 0.95
build_vision_towerFunction · 0.85
build_vision_projectorFunction · 0.85
load_modelMethod · 0.45

Tested by

no test coverage detected