MCPcopy Create free account
hub / github.com/FunAudioLLM/Fun-ASR / __init__

Method __init__

model.py:28–159  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

26@tables.register("model_classes", "FunASRNano")
27class 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):

Callers

nothing calls this directly

Calls 2

CTCClass · 0.90
from_pretrainedMethod · 0.80

Tested by

no test coverage detected