rust-ml / rust-ml/linfa

MultiLogisticRegression panics on normalized data.

Open
#334 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Rust
Stars
4.7k
Forks
334
Avg merge
39m
Merged PRs (30d)
1

Description

Firstly, let me say I'm very new to data science / ML so my understanding / terminology may be wrong. Please bare with me, thanks in advance.

I'm using a relatively small dataset (769 features over 6893 samples) and a small number of categories (8). All of my weights are normalized between [0,1], though most are zero. I'm using the default configuration:

let model = MultiLogisticRegression::default().fit(&dataset).unwrap();
thread 'main' panicked at src/train.rs:121:22:
called `Result::unwrap()` on an `Err` value: ArgMinError(Condition violated: "`MoreThuenteLineSearch`: Search direction must be a descent direction.")

I've observed that if I set alpha to a larger value like 10, I get a result without panicking. I've also noticed that if I limit the number of iterations to a very small number, say, 20, I also get a good result without panicking. Therefore, I think the culprit is overfitting / divergence (uncertain of the proper terminology here).

I will say that I was able to use smartcore's multinomial logistic regression routines while setting alpha = 0 without issue. Notably, they use a Backtracking line search implementation instead of More-Thuente. I don't know if that's relevant or not.

I think this situation is related to this comment on the original MultiLogisticRegression PR regarding divergence. If this is something that can be addressed by linfa by using a different line search algorithm with different numerical requirements, great. If changing the default line search algorithm is undesirable, then at least letting users configure the algorithm used would be greatly appreciated. Most of all however, I would suggest printing a significantly more helpful error message when this divergence happens. If linfa could catch the error returned by argmin and translate it to something along the lines of "your dataset diverged, please increase alpha or reduce the number of iterations", I imagine you would save a ton of developer troubleshooting hours.

Thanks again, and please let me know if there's any other details I should provide here.

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start at src/train.rs:121 and trace how MultiLogisticRegression::fit handles the ArgMinError from MoreThuenteLineSearch. Reproduce the panic with normalized data, then compare behavior when alpha or the iteration limit changes. Done means the divergence is handled or reported more clearly, with any line-search configurability defined by the chosen approach.

Written by the indexing model from the issue text.

Assessment

Tech stack
rust
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
30/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.