(
self,
audio_encoder: str = None,
audio_encoder_conf: dict = None,
audio_adaptor: str = None,
audio_adaptor_conf: dict = None,
llm: str = None,
llm_conf: dict = None,
input_size: int = 80,
length_normalized_loss: bool = False,
**kwargs,
)
| 26 | @tables.register("model_classes", "FunASRNano") |
| 27 | class FunASRNano(nn.Module): |
| 28 | def __init__( |
| 29 | self, |
| 30 | audio_encoder: str = None, |
| 31 | audio_encoder_conf: dict = None, |
| 32 | audio_adaptor: str = None, |
| 33 | audio_adaptor_conf: dict = None, |
| 34 | llm: str = None, |
| 35 | llm_conf: dict = None, |
| 36 | input_size: int = 80, |
| 37 | length_normalized_loss: bool = False, |
| 38 | **kwargs, |
| 39 | ): |
| 40 | super().__init__() |
| 41 | |
| 42 | # audio encoder |
| 43 | hub = audio_encoder_conf.get("hub", None) |
| 44 | self.audio_encoder_activation_checkpoint = audio_encoder_conf.get( |
| 45 | "activation_checkpoint", False |
| 46 | ) |
| 47 | if hub == "ms": |
| 48 | from funasr import AutoModel |
| 49 | |
| 50 | model = AutoModel(model=audio_encoder, model_revision="master") |
| 51 | audio_encoder_output_size = ( |
| 52 | model.model.encoder_output_size |
| 53 | if hasattr(model.model, "encoder_output_size") |
| 54 | else -1 |
| 55 | ) |
| 56 | audio_encoder = ( |
| 57 | model.model.model.encoder if hasattr(model.model, "model") else model.model.encoder |
| 58 | ) |
| 59 | else: |
| 60 | encoder_class = tables.encoder_classes.get(audio_encoder) |
| 61 | audio_encoder = encoder_class(input_size=input_size, **audio_encoder_conf) |
| 62 | audio_encoder_output_size = audio_encoder.output_size() |
| 63 | freeze = audio_encoder_conf.get("freeze", True) |
| 64 | |
| 65 | if freeze: |
| 66 | for _, param in audio_encoder.named_parameters(): |
| 67 | param.requires_grad = False |
| 68 | audio_encoder.eval() |
| 69 | self.audio_encoder = audio_encoder |
| 70 | |
| 71 | # llm |
| 72 | self.llm = None |
| 73 | init_param_path = llm_conf.get("init_param_path", None) |
| 74 | llm_dim = None |
| 75 | |
| 76 | llm_load_kwargs = llm_conf.get("load_kwargs", {}) |
| 77 | config = AutoConfig.from_pretrained(init_param_path) |
| 78 | model = AutoModelForCausalLM.from_config(config, **llm_load_kwargs) |
| 79 | |
| 80 | freeze = llm_conf.get("freeze", True) |
| 81 | if freeze: |
| 82 | for _, param in model.named_parameters(): |
| 83 | param.requires_grad = False |
| 84 | model.eval() |
| 85 | if llm_conf.get("activation_checkpoint", False): |
nothing calls this directly
no test coverage detected