2018-09-04 02:05:57 +00:00
using System ;
using System.Collections.Generic ;
using System.Linq ;
using System.Text ;
using System.Threading.Tasks ;
using System.IO ;
2018-09-29 04:21:44 +00:00
using Bigtree.Algorithm.CRFLite.Utils ;
2018-09-04 02:05:57 +00:00
namespace Txt2Vec
{
public class Term
{
public string strTerm ;
public float [ ] vector ;
public byte [ ] vectorVQ ;
}
public class Model
{
private Dictionary < string , Term > term2vector ;
private List < Term > entireTermList ;
private int vectorSize ;
private double [ ] [ ] codebooks ;
public List < Term > Vocabulary { get { return entireTermList ; } }
public int VectorSize { get { return vectorSize ; } }
public string [ ] GetAllTerms ( )
{
return term2vector . Keys . ToArray ( ) ;
}
public Term GetTerm ( string strTerm )
{
if ( term2vector . ContainsKey ( strTerm ) = = false )
{
return null ;
}
return term2vector [ strTerm ] ;
}
public void LoadModel ( string strFileName , bool bTextFormat )
{
if ( bTextFormat = = true )
{
LoadTextModel ( strFileName ) ;
}
else
{
LoadBinaryModel ( strFileName ) ;
}
}
public bool DumpModel ( string strFileName )
{
if ( entireTermList = = null | | entireTermList . Count = = 0 )
{
return false ;
}
StreamWriter sw = new StreamWriter ( strFileName ) ;
foreach ( Term term in entireTermList )
{
StringBuilder sb = new StringBuilder ( ) ;
sb . Append ( term . strTerm ) ;
sb . Append ( "\t" ) ;
foreach ( double v in term . vector )
{
sb . Append ( v ) ;
sb . Append ( "\t" ) ;
}
sw . WriteLine ( sb . ToString ( ) . Trim ( ) ) ;
}
sw . Close ( ) ;
return true ;
}
public void LoadTextModel ( string strFileName )
{
term2vector = new Dictionary < string , Term > ( ) ;
entireTermList = new List < Term > ( ) ;
vectorSize = 0 ;
StreamReader sr = new StreamReader ( strFileName ) ;
string strLine = null ;
while ( ( strLine = sr . ReadLine ( ) ) ! = null )
{
//the format is "word \t vector
//eah dim of vector is splitted by \t
Term term = new Term ( ) ;
string [ ] items = strLine . Split ( '\t' ) ;
int vSize = items . Length - 1 ;
if ( vectorSize > 0 & & vectorSize ! = vSize )
{
throw new InvalidDataException ( String . Format ( "Invalidated data : {0} . The length of vector must be fixed (current length {1} != previous length {2})." , strLine , vSize , vectorSize ) ) ;
}
term . strTerm = items [ 0 ] ;
term . vector = new float [ vSize ] ;
for ( int i = 0 ; i < vSize ; i + + )
{
term . vector [ i ] = float . Parse ( items [ i + 1 ] ) ;
}
vectorSize = vSize ;
term . vector = NormalizeVector ( term . vector ) ;
term2vector . Add ( term . strTerm , term ) ;
entireTermList . Add ( term ) ;
}
sr . Close ( ) ;
}
public float [ ] GetVector ( string strTerm )
{
if ( term2vector . ContainsKey ( strTerm ) = = true )
{
return term2vector [ strTerm ] . vector ;
}
return null ;
}
private float [ ] NormalizeVector ( float [ ] vec )
{
//Normalize the vector
double len = 0 ;
for ( int a = 0 ; a < vectorSize ; a + + )
{
len + = vec [ a ] * vec [ a ] ;
}
len = Math . Sqrt ( len ) ;
for ( int a = 0 ; a < vectorSize ; a + + )
{
vec [ a ] = ( float ) ( vec [ a ] / len ) ;
}
return vec ;
}
public void LoadBinaryModel ( string strFileName )
{
StreamReader sr = new StreamReader ( strFileName ) ;
BinaryReader br = new BinaryReader ( sr . BaseStream ) ;
//The number of words
int words = br . ReadInt32 ( ) ;
//The size of vector
vectorSize = br . ReadInt32 ( ) ;
int vqSize = br . ReadInt32 ( ) ;
term2vector = new Dictionary < string , Term > ( ) ;
entireTermList = new List < Term > ( ) ;
// Logger.WriteLine("vocabulary size: {0}, vector size: {1}, VQ size: {2}", words, vectorSize, vqSize);
codebooks = null ;
if ( vqSize > 0 )
{
//Read code books
codebooks = new double [ vectorSize ] [ ] ;
for ( int i = 0 ; i < vectorSize ; i + + )
{
codebooks [ i ] = new double [ vqSize ] ;
for ( int j = 0 ; j < vqSize ; j + + )
{
codebooks [ i ] [ j ] = br . ReadDouble ( ) ;
}
}
}
for ( int b = 0 ; b < words ; b + + )
{
Term term = new Term ( ) ;
term . strTerm = br . ReadString ( ) ;
term . vector = new float [ vectorSize ] ;
if ( codebooks ! = null )
{
term . vectorVQ = new byte [ vectorSize ] ;
}
else
{
term . vectorVQ = null ;
}
for ( int i = 0 ; i < vectorSize ; i + + )
{
if ( codebooks = = null )
{
term . vector [ i ] = br . ReadSingle ( ) ;
}
else
{
byte idx = br . ReadByte ( ) ;
term . vector [ i ] = ( float ) codebooks [ i ] [ idx ] ;
term . vectorVQ [ i ] = idx ;
}
}
term . vector = NormalizeVector ( term . vector ) ;
term2vector . Add ( term . strTerm , term ) ;
entireTermList . Add ( term ) ;
}
sr . Close ( ) ;
}
public static void SaveModel ( string strFileName , int vocab_size , int vector_size , List < vocab_word > vocab , double [ ] syn )
{
StreamWriter fo = new StreamWriter ( strFileName ) ;
BinaryWriter bw = new BinaryWriter ( fo . BaseStream ) ;
// Logger.WriteLine("Saving term and vector into model file...");
// Save the word vectors
bw . Write ( vocab_size ) ;
bw . Write ( vector_size ) ;
bw . Write ( 0 ) ; //no VQ
for ( int i = 0 ; i < vocab_size ; i + + )
{
//term string
bw . Write ( vocab [ i ] . word ) ;
//term vector
for ( int j = 0 ; j < vector_size ; j + + )
{
bw . Write ( ( float ) ( syn [ i * vector_size + j ] ) ) ;
}
}
bw . Flush ( ) ;
fo . Flush ( ) ;
fo . Close ( ) ;
}
public bool BuildVQModel ( string strFileName )
{
int vqSize = 256 ;
if ( entireTermList = = null | | entireTermList . Count = = 0 )
{
return false ;
}
StreamWriter fo = new StreamWriter ( strFileName ) ;
BinaryWriter bw = new BinaryWriter ( fo . BaseStream ) ;
// Save the word vectors
bw . Write ( entireTermList . Count ) ; //Vocabulary size
bw . Write ( vectorSize ) ; //Vector size
bw . Write ( vqSize ) ; //VQ size
// Logger.WriteLine("vocabulary size: {0}, vector size: {1}, vq size: {2}", entireTermList.Count, vectorSize, vqSize);
//Create word and VQ values mapping table
Dictionary < string , List < byte > > vqResult = new Dictionary < string , List < byte > > ( ) ;
foreach ( Term term in entireTermList )
{
vqResult . Add ( term . strTerm , new List < byte > ( ) ) ;
}
// Logger.WriteLine("Dims Distortion:");
for ( int i = 0 ; i < vectorSize ; i + + )
{
//Generate VQ values for each dimension
VectorQuantization vq = new VectorQuantization ( ) ;
for ( int j = 0 ; j < entireTermList . Count ; j + + )
{
vq . Add ( entireTermList [ j ] . vector [ i ] ) ;
}
double distortion = vq . BuildCodebook ( vqSize ) ;
// Logger.WriteLine("Dim {0}: {1}", i, distortion);
for ( int j = 0 ; j < entireTermList . Count ; j + + )
{
byte vqValue = ( byte ) vq . ComputeVQ ( entireTermList [ j ] . vector [ i ] ) ;
vqResult [ entireTermList [ j ] . strTerm ] . Add ( vqValue ) ;
}
//Save VQ codebook into model file
for ( int j = 0 ; j < vqSize ; j + + )
{
bw . Write ( vq . CodeBook [ j ] ) ;
}
}
foreach ( KeyValuePair < string , List < byte > > pair in vqResult )
{
if ( pair . Value . Count ! = vectorSize )
{
throw new Exception ( String . Format ( "word {0} has inconsistent vector size: orginial size is {1}, vq size is {2}" ,
pair . Key , vectorSize , pair . Value . Count ) ) ;
}
//term string
bw . Write ( pair . Key ) ;
//term vector
for ( int b = 0 ; b < pair . Value . Count ; b + + )
{
bw . Write ( pair . Value [ b ] ) ;
}
}
bw . Flush ( ) ;
fo . Flush ( ) ;
fo . Close ( ) ;
return true ;
}
}
}