@@ -69,8 +69,8 @@ protected override void TrainCore(Matrix<T> x, Vector<T> y)
6969 {
7070 Vector < T > input = x . GetRow ( i ) ;
7171
72- // Forward pass
73- var ( mean , logVar ) = _encoder . Encode ( input ) ;
72+ // Forward pass - logVar is available for full VAE training with KL loss
73+ var ( mean , _ ) = _encoder . Encode ( input ) ;
7474
7575 // Reparameterization trick (simplified - use mean in deterministic mode)
7676 Vector < T > z = mean . Clone ( ) ;
@@ -237,6 +237,10 @@ protected override void DeserializeCore(BinaryReader reader)
237237 _options . LatentDim = reader . ReadInt32 ( ) ;
238238 _reconstructionThreshold = _numOps . FromDouble ( reader . ReadDouble ( ) ) ;
239239
240+ // Rebuild encoder/decoder to match deserialized dimensions
241+ _encoder = new LSTMEncoder < T > ( _options . WindowSize , _options . LatentDim , _options . HiddenSize ) ;
242+ _decoder = new LSTMDecoder < T > ( _options . LatentDim , _options . WindowSize , _options . HiddenSize ) ;
243+
240244 int encoderParamCount = reader . ReadInt32 ( ) ;
241245 var encoderParams = new Vector < T > ( encoderParamCount ) ;
242246 for ( int i = 0 ; i < encoderParamCount ; i ++ )
@@ -362,7 +366,7 @@ private Matrix<T> CreateRandomMatrix(int rows, int cols, double stddev, Random r
362366 {
363367 sum = _numOps . Add ( sum , _numOps . Multiply ( _weights [ i , j ] , input [ j ] ) ) ;
364368 }
365- hidden [ i ] = _numOps . Tanh ( sum ) ;
369+ hidden [ i ] = MathHelper . Tanh ( sum ) ;
366370 }
367371
368372 // Compute mean
@@ -395,22 +399,44 @@ private Matrix<T> CreateRandomMatrix(int rows, int cols, double stddev, Random r
395399 public Vector < T > GetParameters ( )
396400 {
397401 var parameters = new List < T > ( ) ;
402+ // Include all weights that contribute to ParameterCount
398403 for ( int i = 0 ; i < _weights . Rows ; i ++ )
399404 for ( int j = 0 ; j < _weights . Columns ; j ++ )
400405 parameters . Add ( _weights [ i , j ] ) ;
401406 for ( int i = 0 ; i < _bias . Length ; i ++ )
402407 parameters . Add ( _bias [ i ] ) ;
408+ for ( int i = 0 ; i < _meanWeights . Rows ; i ++ )
409+ for ( int j = 0 ; j < _meanWeights . Columns ; j ++ )
410+ parameters . Add ( _meanWeights [ i , j ] ) ;
411+ for ( int i = 0 ; i < _meanBias . Length ; i ++ )
412+ parameters . Add ( _meanBias [ i ] ) ;
413+ for ( int i = 0 ; i < _logVarWeights . Rows ; i ++ )
414+ for ( int j = 0 ; j < _logVarWeights . Columns ; j ++ )
415+ parameters . Add ( _logVarWeights [ i , j ] ) ;
416+ for ( int i = 0 ; i < _logVarBias . Length ; i ++ )
417+ parameters . Add ( _logVarBias [ i ] ) ;
403418 return new Vector < T > ( parameters . ToArray ( ) ) ;
404419 }
405420
406421 public void SetParameters ( Vector < T > parameters )
407422 {
408423 int idx = 0 ;
424+ // Set all weights that contribute to ParameterCount
409425 for ( int i = 0 ; i < _weights . Rows && idx < parameters . Length ; i ++ )
410426 for ( int j = 0 ; j < _weights . Columns && idx < parameters . Length ; j ++ )
411427 _weights [ i , j ] = parameters [ idx ++ ] ;
412428 for ( int i = 0 ; i < _bias . Length && idx < parameters . Length ; i ++ )
413429 _bias [ i ] = parameters [ idx ++ ] ;
430+ for ( int i = 0 ; i < _meanWeights . Rows && idx < parameters . Length ; i ++ )
431+ for ( int j = 0 ; j < _meanWeights . Columns && idx < parameters . Length ; j ++ )
432+ _meanWeights [ i , j ] = parameters [ idx ++ ] ;
433+ for ( int i = 0 ; i < _meanBias . Length && idx < parameters . Length ; i ++ )
434+ _meanBias [ i ] = parameters [ idx ++ ] ;
435+ for ( int i = 0 ; i < _logVarWeights . Rows && idx < parameters . Length ; i ++ )
436+ for ( int j = 0 ; j < _logVarWeights . Columns && idx < parameters . Length ; j ++ )
437+ _logVarWeights [ i , j ] = parameters [ idx ++ ] ;
438+ for ( int i = 0 ; i < _logVarBias . Length && idx < parameters . Length ; i ++ )
439+ _logVarBias [ i ] = parameters [ idx ++ ] ;
414440 }
415441}
416442
@@ -469,7 +495,7 @@ public Vector<T> Decode(Vector<T> latent)
469495 {
470496 sum = _numOps . Add ( sum , _numOps . Multiply ( _weights [ i , j ] , latent [ j ] ) ) ;
471497 }
472- hidden [ i ] = _numOps . Tanh ( sum ) ;
498+ hidden [ i ] = MathHelper . Tanh ( sum ) ;
473499 }
474500
475501 // Decode to output
@@ -490,21 +516,33 @@ public Vector<T> Decode(Vector<T> latent)
490516 public Vector < T > GetParameters ( )
491517 {
492518 var parameters = new List < T > ( ) ;
519+ // Include all weights that contribute to ParameterCount
493520 for ( int i = 0 ; i < _weights . Rows ; i ++ )
494521 for ( int j = 0 ; j < _weights . Columns ; j ++ )
495522 parameters . Add ( _weights [ i , j ] ) ;
496523 for ( int i = 0 ; i < _bias . Length ; i ++ )
497524 parameters . Add ( _bias [ i ] ) ;
525+ for ( int i = 0 ; i < _outputWeights . Rows ; i ++ )
526+ for ( int j = 0 ; j < _outputWeights . Columns ; j ++ )
527+ parameters . Add ( _outputWeights [ i , j ] ) ;
528+ for ( int i = 0 ; i < _outputBias . Length ; i ++ )
529+ parameters . Add ( _outputBias [ i ] ) ;
498530 return new Vector < T > ( parameters . ToArray ( ) ) ;
499531 }
500532
501533 public void SetParameters ( Vector < T > parameters )
502534 {
503535 int idx = 0 ;
536+ // Set all weights that contribute to ParameterCount
504537 for ( int i = 0 ; i < _weights . Rows && idx < parameters . Length ; i ++ )
505538 for ( int j = 0 ; j < _weights . Columns && idx < parameters . Length ; j ++ )
506539 _weights [ i , j ] = parameters [ idx ++ ] ;
507540 for ( int i = 0 ; i < _bias . Length && idx < parameters . Length ; i ++ )
508541 _bias [ i ] = parameters [ idx ++ ] ;
542+ for ( int i = 0 ; i < _outputWeights . Rows && idx < parameters . Length ; i ++ )
543+ for ( int j = 0 ; j < _outputWeights . Columns && idx < parameters . Length ; j ++ )
544+ _outputWeights [ i , j ] = parameters [ idx ++ ] ;
545+ for ( int i = 0 ; i < _outputBias . Length && idx < parameters . Length ; i ++ )
546+ _outputBias [ i ] = parameters [ idx ++ ] ;
509547 }
510548}
0 commit comments