Finnish-NLP/Ahma-3B
12146
1import json2import mlxu3from EasyLM.serving import LMClient4 5 6FLAGS, FLAGS_DEF = mlxu.define_flags_with_default(7 input_file='',8 output_file='',9 prefix_field='prefix',10 text_field='text',11 until_field='until',12 eval_type='loglikelihood',13 lm_client=LMClient.get_default_config(),14)15 16 17def main(argv):18 lm_client = LMClient(FLAGS.lm_client)19 with mlxu.open_file(FLAGS.input_file, 'r') as fin:20 input_data = json.load(fin)21 22 if FLAGS.eval_type == 'loglikelihood':23 prefix = input_data[FLAGS.prefix_field]24 text = input_data[FLAGS.text_field]25 loglikelihoods, is_greedys = lm_client.loglikelihood(prefix, text)26 output_data = {27 'loglikelihood': loglikelihoods,28 'is_greedy': is_greedys,29 }30 elif FLAGS.eval_type == 'loglikelihood_rolling':31 text = input_data[FLAGS.text_field]32 loglikelihoods, is_greedys = lm_client.loglikelihood_rolling(text)33 output_data = {34 'loglikelihood': loglikelihoods,35 'is_greedy': is_greedys,36 }37 elif FLAGS.eval_type == 'greedy_until':38 prefix = input_data[FLAGS.prefix_field]39 until = input_data[FLAGS.until_field]40 output_data = {'output_text': lm_client.greedy_until(prefix, until)}41 elif FLAGS.eval_type == 'generate':42 prefix = input_data[FLAGS.prefix_field]43 output_data = {'output_text': lm_client.generate(prefix)}44 else:45 raise ValueError(f'Unknown eval_type: {FLAGS.eval_type}')46 47 with mlxu.open_file(FLAGS.output_file, 'w') as fout:48 json.dump(output_data, fout)49 50 51if __name__ == "__main__":52 mlxu.run(main)53 