CoolFace
Apppublic

becaliang/Music_Generation_Project

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
lstm.py77 linesDownload Raw Back to root
1import numpy as np2from keras.models import Sequential3from keras.layers import LSTM, Dense, Dropout, Activation4from keras.utils import to_categorical5 6def _one_hot_encode_sequences(X_data, y_data, token_to_index, sequence_length):7    vocab_size = len(token_to_index)8    X = np.zeros((len(X_data), sequence_length, vocab_size), dtype=np.float32)9    y = np.zeros((len(X_data), vocab_size), dtype=np.float32)10 11    for i, seq in enumerate(X_data):12        for t, token in enumerate(seq):13            if token in token_to_index:14                X[i, t, token_to_index[token]] = 1.015        if y_data[i] in token_to_index:16            y[i, token_to_index[y_data[i]]] = 1.017    return X, y18 19def build_multi_instrument_model(track_corpora, token_to_index, sequence_length=20, N_epochs=64):20    """21    Builds and trains an LSTM model for symbolic multi-instrument music generation.22 23    Parameters:24    - track_corpora: list of token sequences (each list = 1 instrument track)25    - token_to_index: dictionary mapping token -> index26    - sequence_length: number of tokens in each input sequence27    - N_epochs: number of training epochs28    """29    stride = 330    vocab_size = len(token_to_index)31    X_data, y_data = [], []32 33    # === Build training data ===34    for idx, corpus in enumerate(track_corpora):35        track_name = f"Track_{idx + 1}"36        if len(corpus) <= sequence_length:37            print(f"⚠️ Skipping {track_name}: too short ({len(corpus)} tokens).")38            continue39        for i in range(0, len(corpus) - sequence_length, stride):40            X_data.append(corpus[i:i + sequence_length])41            y_data.append(corpus[i + sequence_length])42 43    if len(X_data) == 0:44        # Fallback: pad with dummy "<PAD>" tokens45        print("⚠️ No valid sequences. Padding with dummy data to allow model structure initialization.")46        dummy_token = "<PAD>"47        if dummy_token not in token_to_index:48            token_to_index[dummy_token] = vocab_size49            vocab_size += 150        dummy_seq = [dummy_token] * sequence_length51        X_data = [dummy_seq, dummy_seq]52        y_data = [dummy_token, dummy_token]53 54    print(f"✅ Total sequences: {len(X_data)} | Vocabulary size: {vocab_size}")55 56    # === One-hot encode ===57    X, y = _one_hot_encode_sequences(X_data, y_data, token_to_index, sequence_length)58 59    # === Build LSTM model ===60    model = Sequential([61        LSTM(256, return_sequences=True, input_shape=(sequence_length, vocab_size)),62        Dropout(0.3),63        LSTM(256),64        Dropout(0.3),65        Dense(vocab_size),66        Activation("softmax")67    ])68 69    model.compile(70        loss="categorical_crossentropy",71        optimizer="rmsprop",72        metrics=["accuracy"]73    )74 75    model.fit(X, y, batch_size=64, epochs=N_epochs, verbose=1)76 77    return model